[优化] 目录调整

This commit is contained in:
ryan
2026-06-06 16:08:40 +08:00
parent 314229b7ee
commit 959b134d67
230 changed files with 468 additions and 296 deletions
@@ -0,0 +1,178 @@
package common
import (
"sync"
"time"
"github.com/google/uuid"
)
var StartTime = time.Now().Unix() // unit: second
var Version = "dev" // release builds inject the tag version via ldflags
var SystemName = "OpenFlare"
var ServerAddress = "http://localhost:3000"
var Footer = ""
var HomePageLink = ""
// Any options with "Secret", "Token" in its key won't be return by GetOptions
var SessionSecret = uuid.New().String()
var JWTSecret = "" // if empty, falls back to SessionSecret; set via JWT_SECRET env var
var SQLitePath = "openflare.db"
var SQLDSN = ""
var OptionMap map[string]string
var OptionMapRWMutex sync.RWMutex
var ItemsPerPage = 10
var PasswordLoginEnabled = true
var CapLoginEnabled = true
var PasswordRegisterEnabled = false
var EmailVerificationEnabled = false
var GitHubOAuthEnabled = false
var WeChatAuthEnabled = false
var RegisterEnabled = false
var SMTPServer = ""
var SMTPPort = 587
var SMTPAccount = ""
var SMTPToken = ""
var GitHubClientId = ""
var GitHubClientSecret = ""
var WeChatServerAddress = ""
var WeChatServerToken = ""
var WeChatAccountQRCodeImageURL = ""
var AccessToken = ""
var AgentDiscoveryToken = ""
var NodeOfflineThreshold = 2 * time.Minute
// V3 operational settings (hot-reloadable via Option table)
var AgentHeartbeatInterval = 10000 // milliseconds
var AgentWebsocketUpgradeEnabled = true
var AgentUpdateRepo = "Rain-kl/OpenFlare"
var GeoIPProvider = "ipinfo"
var DatabaseAutoCleanupEnabled = false
var DatabaseAutoCleanupRetentionDays = 30
// Uptime Kuma integration settings
var UptimeKumaEnabled = false
var UptimeKumaUrl = ""
var UptimeKumaUsername = ""
var UptimeKumaPassword = ""
var UptimeKumaMonitorScope = "all" // "all" or "selected"
var UptimeKumaSelectedSites = "" // Comma-separated list of site names
var UptimeKumaSyncInterval = 5 // minutes
var UptimeKumaInterval = 60 // seconds
var UptimeKumaRetry = 0
var UptimeKumaRetryInterval = 60 // seconds
var UptimeKumaTimeout = 48 // seconds
// V5 OpenResty performance settings (hot-reloadable via Option table)
var OpenRestyDefaultServerReturnStatus = 421
var OpenRestyWorkerProcesses = "auto"
var OpenRestyWorkerConnections = 4096
var OpenRestyWorkerRlimitNofile = 65535
var OpenRestyEventsUse = "epoll"
var OpenRestyEventsMultiAcceptEnabled = true
var OpenRestyKeepaliveTimeout = 20
var OpenRestyKeepaliveRequests = 1000
var OpenRestyClientHeaderTimeout = 15
var OpenRestyClientBodyTimeout = 15
var OpenRestyClientMaxBodySize = "64m"
var OpenRestyLargeClientHeaderBuffers = "4 16k"
var OpenRestySendTimeout = 30
var OpenRestyResolvers = ""
var OpenRestyProxyConnectTimeout = 3
var OpenRestyProxySendTimeout = 60
var OpenRestyProxyReadTimeout = 60
var OpenRestyWebsocketEnabled = true
var OpenRestyHTTP3Enabled = true
var OpenRestyProxyRequestBufferingEnabled = false
var OpenRestyProxyBufferingEnabled = true
var OpenRestyProxyBuffers = "16 16k"
var OpenRestyProxyBufferSize = "8k"
var OpenRestyProxyBusyBuffersSize = "64k"
var OpenRestyGzipEnabled = true
var OpenRestyGzipMinLength = 1024
var OpenRestyGzipCompLevel = 5
var OpenRestyCacheEnabled = false
var OpenRestyCachePath = ""
var OpenRestyCacheLevels = "1:2"
var OpenRestyCacheInactive = "30m"
var OpenRestyCacheMaxSize = "1g"
var OpenRestyCacheKeyTemplate = "$scheme$host$request_uri"
var OpenRestyCacheLockEnabled = true
var OpenRestyCacheLockTimeout = "5s"
var OpenRestyCacheUseStale = "error timeout updating http_500 http_502 http_503 http_504"
var OpenRestyMainConfigTemplate = `# This file is generated by OpenFlare. Do not edit manually.
worker_processes {{OpenRestyWorkerProcesses}};
worker_rlimit_nofile {{OpenRestyWorkerRlimitNofile}};
pid logs/nginx.pid;
error_log {{OpenRestyErrorLogPath}} warn;
events {
worker_connections {{OpenRestyWorkerConnections}};
{{OpenRestyEventsUseDirective}}{{OpenRestyEventsMultiAcceptDirective}}}
http {
include mime.types;
default_type application/octet-stream;
{{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}';
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}};
}
`
const (
RoleGuestUser = 0
RoleCommonUser = 1
RoleAdminUser = 10
RoleRootUser = 100
)
// All duration's unit is seconds
// Shouldn't larger then RateLimitKeyExpirationDuration
var (
GlobalApiRateLimitNum = 300
GlobalApiRateLimitDuration int64 = 3 * 60
GlobalWebRateLimitNum = 300
GlobalWebRateLimitDuration int64 = 3 * 60
CriticalRateLimitNum = 100
CriticalRateLimitDuration int64 = 20 * 60
)
var RateLimitKeyExpirationDuration = 20 * time.Minute
const (
UserStatusEnabled = 1 // don't use 0, 0 is the default value!
UserStatusDisabled = 2 // also don't use 0
)
+84
View File
@@ -0,0 +1,84 @@
package common
import (
"flag"
"fmt"
"log/slog"
"os"
"path/filepath"
"strings"
)
var (
Port = flag.Int("port", 3000, "the listening port")
PrintVersion = flag.Bool("version", false, "print version and exit")
PrintHelp = flag.Bool("help", false, "print help and exit")
LogDir = flag.String("log-dir", "", "specify the log directory")
)
func printHelp() {
fmt.Println("OpenFlare " + Version + " - Internal OpenResty Control Plane.")
fmt.Println("Copyright (C) 2023 JustSong. All rights reserved.")
fmt.Println("GitHub: https://github.com/Rain-kl/OpenFlare")
fmt.Println("Usage: openflare [--port <port>] [--log-dir <log directory>] [--version] [--help]")
}
// ParseFlags 在命令行参数被任何 import 链上的 init() 误解析之前,
// 由各 binary 的 main() 显式调用一次。openflare-server 与 openflare-relay
// 共用 flag.CommandLine,必须先注册各自的 flag 再调用本函数。
// 测试场景(go test)下不会执行本函数,单元测试可直接跳过命令行解析。
func ParseFlags() {
executableName := strings.ToLower(filepath.Base(os.Args[0]))
isTest := strings.Contains(executableName, ".test") || flag.Lookup("test.v") != nil
if isTest {
return
}
flag.Parse()
if *PrintVersion {
fmt.Println(Version)
os.Exit(0)
}
if *PrintHelp {
printHelp()
os.Exit(0)
}
if os.Getenv("SESSION_SECRET") != "" {
SessionSecret = os.Getenv("SESSION_SECRET")
}
if os.Getenv("JWT_SECRET") != "" {
JWTSecret = os.Getenv("JWT_SECRET")
}
if os.Getenv("SQLITE_PATH") != "" {
SQLitePath = os.Getenv("SQLITE_PATH")
}
if os.Getenv("SQL_DSN") != "" {
SQLDSN = os.Getenv("SQL_DSN")
}
if os.Getenv("DSN") != "" {
SQLDSN = os.Getenv("DSN")
}
if os.Getenv("AGENT_TOKEN") != "" {
AccessToken = os.Getenv("AGENT_TOKEN")
}
SetLogLevel(os.Getenv("LOG_LEVEL"))
if *LogDir != "" {
var err error
*LogDir, err = filepath.Abs(*LogDir)
if err != nil {
slog.Error("resolve log directory failed", "error", err)
os.Exit(1)
}
if _, err := os.Stat(*LogDir); os.IsNotExist(err) {
err = os.Mkdir(*LogDir, 0777)
if err != nil {
slog.Error("create log directory failed", "error", err)
os.Exit(1)
}
}
}
}
+242
View File
@@ -0,0 +1,242 @@
package common
import (
"context"
"fmt"
"io"
"log/slog"
"os"
"path/filepath"
"runtime"
"slices"
"strings"
"github.com/gin-gonic/gin"
)
type logLevel int
const (
logLevelDebug logLevel = iota
logLevelInfo
logLevelWarn
logLevelError
)
var currentLogLevel = logLevelInfo
var currentLogLevelName = "info"
var commonLogWriter io.Writer = os.Stdout
var errorLogWriter io.Writer = os.Stderr
var defaultLogger *slog.Logger
type customTextHandler struct {
writer io.Writer
level slog.Level
attrs []slog.Attr
groups []string
}
type levelRouterHandler struct {
commonHandler slog.Handler
errorHandler slog.Handler
}
func (h *customTextHandler) Enabled(_ context.Context, level slog.Level) bool {
return level >= h.level
}
func (h *customTextHandler) Handle(_ context.Context, record slog.Record) error {
var builder strings.Builder
builder.WriteString(record.Time.Format("2006-01-02 15:04:05.000"))
builder.WriteString(" | ")
builder.WriteString(fmt.Sprintf("%-8s", levelLabel(record.Level)))
builder.WriteString(" | ")
builder.WriteString(sourceLocation(record.PC))
builder.WriteString(" - ")
builder.WriteString(record.Message)
attrs := make([]slog.Attr, 0, len(h.attrs)+record.NumAttrs())
attrs = append(attrs, h.attrs...)
record.Attrs(func(attr slog.Attr) bool {
attrs = append(attrs, attr)
return true
})
if len(attrs) > 0 {
builder.WriteString(" | ")
builder.WriteString(formatAttrs(h.groups, attrs))
}
builder.WriteByte('\n')
_, err := io.WriteString(h.writer, builder.String())
return err
}
func (h *customTextHandler) WithAttrs(attrs []slog.Attr) slog.Handler {
cloned := *h
cloned.attrs = append(slices.Clone(h.attrs), attrs...)
return &cloned
}
func (h *customTextHandler) WithGroup(name string) slog.Handler {
if strings.TrimSpace(name) == "" {
return h
}
cloned := *h
cloned.groups = append(slices.Clone(h.groups), name)
return &cloned
}
func (h *levelRouterHandler) Enabled(ctx context.Context, level slog.Level) bool {
return h.commonHandler.Enabled(ctx, level) || h.errorHandler.Enabled(ctx, level)
}
func (h *levelRouterHandler) Handle(ctx context.Context, record slog.Record) error {
if record.Level >= slog.LevelError {
return h.errorHandler.Handle(ctx, record)
}
return h.commonHandler.Handle(ctx, record)
}
func (h *levelRouterHandler) WithAttrs(attrs []slog.Attr) slog.Handler {
return &levelRouterHandler{
commonHandler: h.commonHandler.WithAttrs(attrs),
errorHandler: h.errorHandler.WithAttrs(attrs),
}
}
func (h *levelRouterHandler) WithGroup(name string) slog.Handler {
return &levelRouterHandler{
commonHandler: h.commonHandler.WithGroup(name),
errorHandler: h.errorHandler.WithGroup(name),
}
}
func configureGinWriters() {
if shouldLog(logLevelDebug) {
gin.DefaultWriter = commonLogWriter
} else {
gin.DefaultWriter = io.Discard
}
gin.DefaultErrorWriter = errorLogWriter
}
func slogLevel() slog.Level {
switch currentLogLevel {
case logLevelDebug:
return slog.LevelDebug
case logLevelWarn:
return slog.LevelWarn
case logLevelError:
return slog.LevelError
default:
return slog.LevelInfo
}
}
func ensureLogger() *slog.Logger {
if defaultLogger != nil {
return defaultLogger
}
defaultLogger = slog.New(&levelRouterHandler{
commonHandler: &customTextHandler{writer: commonLogWriter, level: slogLevel()},
errorHandler: &customTextHandler{writer: errorLogWriter, level: slogLevel()},
})
slog.SetDefault(defaultLogger)
return defaultLogger
}
func SetLogLevel(level string) {
normalized := strings.TrimSpace(strings.ToLower(level))
switch normalized {
case "debug":
currentLogLevel = logLevelDebug
currentLogLevelName = "debug"
case "warn", "warning":
currentLogLevel = logLevelWarn
currentLogLevelName = "warn"
case "error":
currentLogLevel = logLevelError
currentLogLevelName = "error"
default:
currentLogLevel = logLevelInfo
currentLogLevelName = "info"
}
configureGinWriters()
}
func GetLogLevel() string {
return currentLogLevelName
}
func shouldLog(level logLevel) bool {
return level >= currentLogLevel
}
func SetupGinLog() {
if *LogDir != "" {
commonLogPath := filepath.Join(*LogDir, "common.log")
errorLogPath := filepath.Join(*LogDir, "error.log")
commonFd, err := os.OpenFile(commonLogPath, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644)
if err != nil {
_, _ = io.WriteString(os.Stderr, "failed to open common log file\n")
os.Exit(1)
}
errorFd, err := os.OpenFile(errorLogPath, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644)
if err != nil {
_, _ = io.WriteString(os.Stderr, "failed to open error log file\n")
os.Exit(1)
}
commonLogWriter = io.MultiWriter(os.Stdout, commonFd)
errorLogWriter = io.MultiWriter(os.Stderr, errorFd)
}
configureGinWriters()
defaultLogger = nil
ensureLogger()
}
func levelLabel(level slog.Level) string {
switch {
case level <= slog.LevelDebug:
return "DEBUG"
case level < slog.LevelWarn:
return "INFO"
case level < slog.LevelError:
return "WARNING"
default:
return "ERROR"
}
}
func sourceLocation(pc uintptr) string {
if pc == 0 {
return "unknown:unknown:0"
}
frame, _ := runtime.CallersFrames([]uintptr{pc}).Next()
fileName := strings.TrimSuffix(filepath.Base(frame.File), filepath.Ext(frame.File))
if fileName == "" {
fileName = "unknown"
}
functionName := "unknown"
if frame.Function != "" {
parts := strings.Split(frame.Function, "/")
functionName = parts[len(parts)-1]
if dot := strings.LastIndex(functionName, "."); dot >= 0 && dot < len(functionName)-1 {
functionName = functionName[dot+1:]
}
}
return fmt.Sprintf("%s:%s:%d", fileName, functionName, frame.Line)
}
func formatAttrs(groups []string, attrs []slog.Attr) string {
parts := make([]string, 0, len(attrs))
for _, attr := range attrs {
key := attr.Key
if key == "" {
continue
}
if len(groups) > 0 {
key = strings.Join(append(slices.Clone(groups), key), ".")
}
parts = append(parts, fmt.Sprintf("%s=%v", key, attr.Value.Any()))
}
return strings.Join(parts, " ")
}
+41
View File
@@ -0,0 +1,41 @@
package common
import (
"context"
"log/slog"
"os"
"time"
"github.com/go-redis/redis/v8"
)
var RDB *redis.Client
var RedisEnabled = true
// InitRedisClient This function is called after init()
func InitRedisClient() (err error) {
if os.Getenv("REDIS_CONN_STRING") == "" {
RedisEnabled = false
slog.Info("redis disabled because REDIS_CONN_STRING is not set")
return nil
}
opt, err := redis.ParseURL(os.Getenv("REDIS_CONN_STRING"))
if err != nil {
panic(err)
}
RDB = redis.NewClient(opt)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_, err = RDB.Ping(ctx).Result()
return err
}
func ParseRedisOption() *redis.Options {
opt, err := redis.ParseURL(os.Getenv("REDIS_CONN_STRING"))
if err != nil {
panic(err)
}
return opt
}
@@ -0,0 +1,82 @@
package response
import (
"net/http"
"github.com/gin-gonic/gin"
)
const invalidParamsMessage = "参数错误"
// RespondSuccess sends a successful response with data
func RespondSuccess(c *gin.Context, data any) {
c.JSON(http.StatusOK, gin.H{
"success": true,
"message": "",
"data": data,
})
}
// RespondSuccessWithExtras sends a successful response with data and extra fields
func RespondSuccessWithExtras(c *gin.Context, data any, extras gin.H) {
payload := gin.H{
"success": true,
"message": "",
"data": data,
}
for key, value := range extras {
payload[key] = value
}
c.JSON(http.StatusOK, payload)
}
// RespondSuccessMessage sends a successful response with a custom message
func RespondSuccessMessage(c *gin.Context, message string) {
c.JSON(http.StatusOK, gin.H{
"success": true,
"message": message,
})
}
// RespondFailure sends a failed response with http.StatusOK and a failure message
func RespondFailure(c *gin.Context, message string) {
c.JSON(http.StatusOK, gin.H{
"success": false,
"message": message,
})
}
// RespondBadRequest sends a bad request response (400)
func RespondBadRequest(c *gin.Context, message string) {
if message == "" {
message = invalidParamsMessage
}
c.JSON(http.StatusBadRequest, gin.H{
"success": false,
"message": message,
})
}
// RespondUnauthorized sends an unauthorized response (401)
func RespondUnauthorized(c *gin.Context, message string) {
c.JSON(http.StatusUnauthorized, gin.H{
"success": false,
"message": message,
})
}
// RespondForbidden sends a forbidden response (403)
func RespondForbidden(c *gin.Context, message string) {
c.JSON(http.StatusForbidden, gin.H{
"success": false,
"message": message,
})
}
// RespondErrorWithStatus sends a response with target HTTP status code and a message
func RespondErrorWithStatus(c *gin.Context, code int, message string) {
c.JSON(code, gin.H{
"success": false,
"message": message,
})
}
@@ -0,0 +1,196 @@
package controller
import (
"strconv"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
)
// GetAccessLogs godoc
// @Summary List access logs
// @Tags AccessLogs
// @Produce json
// @Security OpenFlareTokenAuth
// @Param node_id query string false "Node ID"
// @Param remote_addr query string false "Remote address"
// @Param host query string false "Host"
// @Param path query string false "Path"
// @Param p query int false "Page index"
// @Param page_size query int false "Page size"
// @Param sort_by query string false "Sort by"
// @Param sort_order query string false "Sort order"
// @Success 200 {object} map[string]interface{}
// @Router /api/access-logs/ [get]
func GetAccessLogs(c *gin.Context) {
logs, err := service.ListAccessLogs(readAccessLogQuery(c))
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, logs)
}
// GetFoldedAccessLogs godoc
// @Summary List folded access logs
// @Tags AccessLogs
// @Produce json
// @Security OpenFlareTokenAuth
// @Param node_id query string false "Node ID"
// @Param remote_addr query string false "Remote address"
// @Param host query string false "Host"
// @Param path query string false "Path"
// @Param p query int false "Page index"
// @Param page_size query int false "Page size"
// @Param sort_by query string false "Sort by"
// @Param sort_order query string false "Sort order"
// @Param fold_minutes query int false "Fold minutes"
// @Success 200 {object} map[string]interface{}
// @Router /api/access-logs/folds [get]
func GetFoldedAccessLogs(c *gin.Context) {
query := readAccessLogQuery(c)
query.FoldMinutes = readQueryInt(c, "fold_minutes")
logs, err := service.ListFoldedAccessLogs(query)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, logs)
}
// GetFoldedAccessLogIPs godoc
// @Summary List folded access log IP summaries
// @Tags AccessLogs
// @Produce json
// @Security OpenFlareTokenAuth
// @Param node_id query string false "Node ID"
// @Param remote_addr query string false "Remote address"
// @Param host query string false "Host"
// @Param path query string false "Path"
// @Param bucket_started_at query string true "Bucket started at"
// @Param fold_minutes query int true "Fold minutes"
// @Param p query int false "Page index"
// @Param page_size query int false "Page size"
// @Param sort_by query string false "Sort by"
// @Param sort_order query string false "Sort order"
// @Success 200 {object} map[string]interface{}
// @Router /api/access-logs/folds/ip-summary [get]
func GetFoldedAccessLogIPs(c *gin.Context) {
result, err := service.ListFoldedAccessLogIPs(service.FoldedAccessLogIPQuery{
NodeID: c.Query("node_id"),
RemoteAddr: c.Query("remote_addr"),
Host: c.Query("host"),
Path: c.Query("path"),
BucketStartedAt: c.Query("bucket_started_at"),
FoldMinutes: readQueryInt(c, "fold_minutes"),
Page: readQueryInt(c, "p"),
PageSize: readQueryInt(c, "page_size"),
SortBy: c.Query("sort_by"),
SortOrder: c.Query("sort_order"),
})
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, result)
}
// GetAccessLogIPSummaries godoc
// @Summary List access log IP summaries
// @Tags AccessLogs
// @Produce json
// @Security OpenFlareTokenAuth
// @Param node_id query string false "Node ID"
// @Param remote_addr query string false "Remote address"
// @Param host query string false "Host"
// @Param p query int false "Page index"
// @Param page_size query int false "Page size"
// @Param sort_by query string false "Sort by"
// @Param sort_order query string false "Sort order"
// @Success 200 {object} map[string]interface{}
// @Router /api/access-logs/ip-summary [get]
func GetAccessLogIPSummaries(c *gin.Context) {
result, err := service.ListAccessLogIPSummaries(service.AccessLogIPSummaryQuery{
NodeID: c.Query("node_id"),
RemoteAddr: c.Query("remote_addr"),
Host: c.Query("host"),
Page: readQueryInt(c, "p"),
PageSize: readQueryInt(c, "page_size"),
SortBy: c.Query("sort_by"),
SortOrder: c.Query("sort_order"),
})
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, result)
}
// GetAccessLogIPTrend godoc
// @Summary Get access log IP trend
// @Tags AccessLogs
// @Produce json
// @Security OpenFlareTokenAuth
// @Param node_id query string false "Node ID"
// @Param remote_addr query string true "Remote address"
// @Param host query string false "Host"
// @Param hours query int false "Hours"
// @Param bucket_minutes query int false "Bucket minutes"
// @Success 200 {object} map[string]interface{}
// @Router /api/access-logs/ip-summary/trend [get]
func GetAccessLogIPTrend(c *gin.Context) {
result, err := service.GetAccessLogIPTrend(service.AccessLogIPTrendQuery{
NodeID: c.Query("node_id"),
RemoteAddr: c.Query("remote_addr"),
Host: c.Query("host"),
Hours: readQueryInt(c, "hours"),
BucketMinutes: readQueryInt(c, "bucket_minutes"),
})
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, result)
}
// CleanupAccessLogs godoc
// @Summary Cleanup access logs by retention days
// @Tags AccessLogs
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/access-logs/cleanup [post]
func CleanupAccessLogs(c *gin.Context) {
var input service.AccessLogCleanupInput
if !bind.JSON(c, &input) {
return
}
result, err := service.CleanupAccessLogs(input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, result)
}
func readAccessLogQuery(c *gin.Context) service.AccessLogQuery {
return service.AccessLogQuery{
NodeID: c.Query("node_id"),
RemoteAddr: c.Query("remote_addr"),
Host: c.Query("host"),
Path: c.Query("path"),
Page: readQueryInt(c, "p"),
PageSize: readQueryInt(c, "page_size"),
SortBy: c.Query("sort_by"),
SortOrder: c.Query("sort_order"),
}
}
func readQueryInt(c *gin.Context, key string) int {
value, _ := strconv.Atoi(c.DefaultQuery(key, "0"))
return value
}
@@ -0,0 +1,24 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/gin-gonic/gin"
)
// GetDefaultAcmeAccount godoc
// @Summary Get default ACME account
// @Tags AcmeAccounts
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/acme-accounts/default [get]
func GetDefaultAcmeAccount(c *gin.Context) {
account, err := model.GetDefaultAcmeAccount()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, account)
}
@@ -0,0 +1,363 @@
package controller
import (
"encoding/json"
"log/slog"
"net"
"strconv"
"time"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
"golang.org/x/net/websocket"
)
// AgentRegister godoc
// @Summary Register or discover agent node
// @Tags Agent
// @Accept json
// @Produce json
// @Security AccessTokenAuth
// @Param payload body service.AgentNodePayload true "Agent node payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/agent/nodes/register [post]
func AgentRegister(c *gin.Context) {
var payload service.AgentNodePayload
if !bind.JSON(c, &payload) {
return
}
payload.IP = service.ResolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
var (
result *service.AgentRegistrationResponse
err error
)
if authNode, ok := c.Get("agent_node"); ok {
result, err = service.RegisterNodeWithAccessToken(authNode.(*model.Node), payload)
} else {
result, err = service.RegisterNodeWithDiscovery(payload)
}
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, result)
}
// AgentHeartbeat godoc
// @Summary Report agent heartbeat
// @Tags Agent
// @Accept json
// @Produce json
// @Security AccessTokenAuth
// @Param payload body service.AgentNodePayload true "Agent heartbeat payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/agent/nodes/heartbeat [post]
func AgentHeartbeat(c *gin.Context) {
var payload service.AgentNodePayload
if !bind.JSON(c, &payload) {
return
}
payload.IP = service.ResolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
authNode, ok := c.Get("agent_node")
if !ok {
response.RespondUnauthorized(c, "鏃犳潈杩涜姝ゆ搷浣滐紝Agent Token 鏃犳晥")
return
}
node, err := service.HeartbeatNode(authNode.(*model.Node), payload)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessWithExtras(c, node.Node, gin.H{
"agent_settings": node.AgentSettings,
"active_config": node.ActiveConfig,
"waf_ip_groups": node.WAFIPGroups,
})
}
// AgentSyncWAFIPGroups godoc
// @Summary Sync WAF IP groups for agent
// @Tags Agent
// @Accept json
// @Produce json
// @Security AccessTokenAuth
// @Param payload body service.AgentWAFIPGroupSyncInput true "WAF IP group sync payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/agent/waf/ip-groups/sync [post]
func AgentSyncWAFIPGroups(c *gin.Context) {
var input service.AgentWAFIPGroupSyncInput
if !bind.JSON(c, &input) {
return
}
result, err := service.SyncWAFIPGroupsForAgent(input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, result)
}
// AgentGetActiveConfig godoc
// @Summary Get active config for agent
// @Tags Agent
// @Produce json
// @Security AccessTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/agent/config-versions/active [get]
func AgentGetActiveConfig(c *gin.Context) {
authNode, ok := c.Get("agent_node")
if !ok {
response.RespondUnauthorized(c, "Node object missing from context")
return
}
node := authNode.(*model.Node)
if node.NodeType == "tunnel_client" {
config, err := service.GetFlaredTunnelConfig(node)
if err != nil {
response.RespondFailure(c, "无法生成隧道配置: "+err.Error())
return
}
response.RespondSuccess(c, config)
return
}
config, err := service.GetActiveConfigForAgent()
if err != nil {
response.RespondFailure(c, "当前没有激活版本")
return
}
response.RespondSuccess(c, config)
}
// AgentReportApplyLog godoc
// @Summary Report agent apply result
// @Tags Agent
// @Accept json
// @Produce json
// @Security AccessTokenAuth
// @Param payload body service.ApplyLogPayload true "Apply log payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/agent/apply-logs [post]
func AgentReportApplyLog(c *gin.Context) {
var payload service.ApplyLogPayload
if !bind.JSON(c, &payload) {
return
}
if authNode, ok := c.Get("agent_node"); ok {
payload.NodeID = authNode.(*model.Node).NodeID
}
log, err := service.ReportApplyLog(payload)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, log)
}
// AgentWebSocket godoc
// @Summary Upgrade agent connection to websocket
// @Tags Agent
// @Security AccessTokenAuth
// @Router /api/agent/ws [get]
func AgentWebSocket(c *gin.Context) {
authNode, ok := c.Get("agent_node")
if !ok {
response.RespondUnauthorized(c, "无权进行此操作,Agent Token 无效")
return
}
node := authNode.(*model.Node)
slog.Debug("agent ws upgrade requested", "node_id", node.NodeID, "remote", c.Request.RemoteAddr)
websocket.Handler(func(conn *websocket.Conn) {
client := service.RegisterAgentWSClient(node.NodeID)
defer service.UnregisterAgentWSClient(client)
defer func() {
_ = conn.Close()
slog.Debug("agent ws connection closed", "node_id", node.NodeID)
}()
slog.Debug("agent ws upgrade succeeded", "node_id", node.NodeID, "remote", c.Request.RemoteAddr)
go func() {
<-client.Done()
_ = conn.Close()
}()
go streamAgentWSMessages(c, conn, client)
for {
var message service.AgentWSInboundMessage
_ = conn.SetReadDeadline(time.Now().Add(agentWSReadTimeout()))
if err := websocket.JSON.Receive(conn, &message); err != nil {
if netErr, ok := err.(net.Error); ok && netErr.Timeout() {
slog.Debug("agent ws receive timeout waiting for status or pong", "node_id", node.NodeID, "timeout", agentWSReadTimeout())
return
}
slog.Debug("agent ws receive failed", "node_id", node.NodeID, "error", err)
return
}
slog.Debug("agent ws message received", "node_id", node.NodeID, "type", message.Type)
switch message.Type {
case service.AgentWSMessageTypeStatus:
handleAgentWSStatus(c, node, message)
case service.AgentWSMessageTypePing:
if !service.SendAgentWSPong(node.NodeID) {
slog.Debug("agent ws pong enqueue failed", "node_id", node.NodeID)
}
case service.AgentWSMessageTypePong:
slog.Debug("agent ws pong received", "node_id", node.NodeID)
default:
slog.Debug("agent ws unsupported message type", "node_id", node.NodeID, "type", message.Type)
}
}
}).ServeHTTP(c.Writer, c.Request)
}
func agentWSReadTimeout() time.Duration {
timeout := time.Duration(common.AgentHeartbeatInterval) * time.Millisecond * 3
if timeout < 30*time.Second {
return 30 * time.Second
}
return timeout
}
func agentWSWriteTimeout() time.Duration {
return 10 * time.Second
}
func streamAgentWSMessages(c *gin.Context, conn *websocket.Conn, client *service.WSClient) {
for {
select {
case <-c.Request.Context().Done():
return
case <-client.Done():
return
case message, ok := <-client.Messages():
if !ok {
return
}
_ = conn.SetWriteDeadline(time.Now().Add(agentWSWriteTimeout()))
if err := websocket.JSON.Send(conn, message); err != nil {
slog.Debug("agent ws send failed", "node_id", client.ID(), "error", err)
return
}
}
}
}
func handleAgentWSStatus(c *gin.Context, node *model.Node, message service.AgentWSInboundMessage) {
var payload service.AgentNodePayload
if err := json.Unmarshal(message.Payload, &payload); err != nil {
slog.Debug("agent ws status payload decode failed", "node_id", node.NodeID, "error", err)
return
}
freshNode, err := model.GetNodeByNodeID(node.NodeID)
if err != nil {
slog.Debug("agent ws status reload node failed", "node_id", node.NodeID, "error", err)
return
}
payload.IP = service.ResolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
res, err := service.HeartbeatNode(freshNode, payload)
if err != nil {
slog.Debug("agent ws status handling failed", "node_id", node.NodeID, "error", err)
return
}
settingsSent := service.SendAgentWSSettings(node.NodeID, res.AgentSettings)
activeConfigSent := false
if res.ActiveConfig != nil {
activeConfigSent = service.SendAgentWSActiveConfig(node.NodeID, res.ActiveConfig)
}
wafIPGroupsSent := false
if len(res.WAFIPGroups) > 0 {
wafIPGroupsSent = service.SendAgentWSWAFIPGroups(node.NodeID, res.WAFIPGroups)
}
slog.Debug("agent ws status processed",
"node_id", node.NodeID,
"current_version", payload.CurrentVersion,
"openresty_status", payload.OpenrestyStatus,
"settings_sent", settingsSent,
"active_config_sent", activeConfigSent,
"waf_ip_groups_sent", wafIPGroupsSent,
)
}
// GetNodes godoc
// @Summary List nodes
// @Tags Nodes
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/nodes/ [get]
func GetNodes(c *gin.Context) {
nodes, err := service.ListNodeViews()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, nodes)
}
// GetApplyLogs godoc
// @Summary List apply logs
// @Tags ApplyLogs
// @Produce json
// @Security OpenFlareTokenAuth
// @Param node_id query string false "Node ID"
// @Success 200 {object} map[string]interface{}
// @Router /api/apply-logs/ [get]
func GetApplyLogs(c *gin.Context) {
logs, err := service.ListApplyLogsPage(service.ApplyLogListQuery{
NodeID: c.Query("node_id"),
PageNo: readIntQueryFallback(c, "pageNo", "page_no"),
PageSize: readIntQueryFallback(c, "pageSize", "page_size"),
})
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, logs)
}
// CleanupApplyLogs godoc
// @Summary Cleanup apply logs
// @Tags ApplyLogs
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/apply-logs/cleanup [post]
func CleanupApplyLogs(c *gin.Context) {
var input service.ApplyLogCleanupInput
if !bind.JSON(c, &input) {
return
}
result, err := service.CleanupApplyLogs(input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, result)
}
func readIntQueryFallback(c *gin.Context, primary string, secondary string) int {
value := c.Query(primary)
if value == "" {
value = c.Query(secondary)
}
parsed, _ := strconv.Atoi(value)
return parsed
}
@@ -0,0 +1,358 @@
package controller
import (
"encoding/json"
"fmt"
"net/url"
"strconv"
"strings"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
const pendingExternalAccountSessionKey = "pending_external_account"
type authSourceTogglePayload struct {
IsActive bool `json:"is_active"`
}
type authSourcePayload struct {
Name string `json:"name"`
Type string `json:"type"`
DisplayName string `json:"display_name"`
IsActive bool `json:"is_active"`
ClientID string `json:"client_id"`
ClientSecret string `json:"client_secret"`
OpenIDDiscoveryURL string `json:"openid_discovery_url"`
Scopes string `json:"scopes"`
IconURL string `json:"icon_url"`
}
func (payload authSourcePayload) toModel() model.AuthSource {
return model.AuthSource{
Name: payload.Name,
Type: payload.Type,
DisplayName: payload.DisplayName,
IsActive: payload.IsActive,
ClientID: payload.ClientID,
ClientSecret: payload.ClientSecret,
OpenIDDiscoveryURL: payload.OpenIDDiscoveryURL,
Scopes: payload.Scopes,
IconURL: payload.IconURL,
}
}
func ListAuthSources(c *gin.Context) {
sources, err := model.GetAuthSources()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, sources)
}
func CreateAuthSource(c *gin.Context) {
var payload authSourcePayload
if err := bind.DecodeJSONBody(c.Request.Body, &payload); err != nil {
response.RespondBadRequest(c, "无效的参数")
return
}
source := payload.toModel()
if err := model.CreateAuthSource(&source); err != nil {
response.RespondFailure(c, err.Error())
return
}
source.Sanitize()
response.RespondSuccess(c, source)
}
func UpdateAuthSource(c *gin.Context) {
id, err := parseAuthSourceID(c)
if err != nil {
response.RespondBadRequest(c, err.Error())
return
}
var payload authSourcePayload
if err := bind.DecodeJSONBody(c.Request.Body, &payload); err != nil {
response.RespondBadRequest(c, "无效的参数")
return
}
source := payload.toModel()
source.ID = id
keepSecret := strings.TrimSpace(source.ClientSecret) == ""
if err := model.UpdateAuthSource(&source, keepSecret); err != nil {
response.RespondFailure(c, err.Error())
return
}
updated, err := model.GetAuthSourceByID(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
updated.Sanitize()
response.RespondSuccess(c, updated)
}
func DeleteAuthSource(c *gin.Context) {
id, err := parseAuthSourceID(c)
if err != nil {
response.RespondBadRequest(c, err.Error())
return
}
if err := model.DeleteAuthSource(id); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
func ToggleAuthSource(c *gin.Context) {
id, err := parseAuthSourceID(c)
if err != nil {
response.RespondBadRequest(c, err.Error())
return
}
var payload authSourceTogglePayload
if err := bind.DecodeJSONBody(c.Request.Body, &payload); err != nil {
response.RespondBadRequest(c, "无效的参数")
return
}
if err := model.ToggleAuthSource(id, payload.IsActive); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
func OAuthAuthorize(c *gin.Context) {
source, err := getAuthSourceFromRoute(c)
if err != nil {
response.RespondBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.RespondFailure(c, "认证源未启用")
return
}
if err := source.Validate(); err != nil {
response.RespondFailure(c, err.Error())
return
}
state, err := service.GenerateOAuthState()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
session := sessions.Default(c)
session.Set(oauthStateSessionKey(source.ID), state)
if err := session.Save(); err != nil {
response.RespondFailure(c, "无法保存授权状态,请重试")
return
}
redirectURL := oauthFrontendCallbackURL(c, source.ID)
authorizeURL, err := service.BuildAuthorizeURL(c.Request.Context(), source, redirectURL, state)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, gin.H{"authorize_url": authorizeURL})
}
func OAuthCallback(c *gin.Context) {
source, err := getAuthSourceFromRoute(c)
if err != nil {
response.RespondBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.RespondFailure(c, "认证源未启用")
return
}
session := sessions.Default(c)
expectedState, _ := session.Get(oauthStateSessionKey(source.ID)).(string)
state := c.Query("state")
if expectedState == "" || state == "" || state != expectedState {
response.RespondFailure(c, "授权状态无效,请重新登录")
return
}
session.Delete(oauthStateSessionKey(source.ID))
if err := session.Save(); err != nil {
response.RespondFailure(c, "无法更新授权状态,请重试")
return
}
if oauthError := c.Query("error"); oauthError != "" {
description := c.Query("error_description")
if description == "" {
description = oauthError
}
response.RespondFailure(c, description)
return
}
profile, err := service.ExchangeOAuthProfile(c.Request.Context(), source, c.Query("code"), oauthFrontendCallbackURL(c, source.ID))
if err != nil {
response.RespondFailure(c, err.Error())
return
}
var currentUserID *int
if currentUser := currentUserFromOpenFlareToken(c); currentUser != nil {
currentUserID = &currentUser.Id
}
result, pending, err := service.CompleteOAuthLogin(source, profile, currentUserID)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
if pending != nil {
raw, err := json.Marshal(pending)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
session.Set(pendingExternalAccountSessionKey, string(raw))
if err := session.Save(); err != nil {
response.RespondFailure(c, "无法保存待绑定账号,请重试")
return
}
response.RespondSuccess(c, result)
return
}
if result.User != nil {
cleanUser, err := setLoginToken(result.User)
if err != nil {
response.RespondFailure(c, "无法保存会话信息,请重试")
return
}
result.User = cleanUser
}
response.RespondSuccess(c, result)
}
func LinkExistingOAuthAccount(c *gin.Context) {
session := sessions.Default(c)
raw, _ := session.Get(pendingExternalAccountSessionKey).(string)
if raw == "" {
response.RespondFailure(c, "待绑定第三方账号已失效,请重新登录")
return
}
var pending service.PendingExternalAccount
if err := json.Unmarshal([]byte(raw), &pending); err != nil {
response.RespondFailure(c, "待绑定第三方账号无效,请重新登录")
return
}
var input service.LinkExistingRequest
if err := bind.DecodeJSONBody(c.Request.Body, &input); err != nil {
response.RespondBadRequest(c, "无效的参数")
return
}
user, err := service.LinkPendingExternalAccount(&pending, input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
session.Delete(pendingExternalAccountSessionKey)
if err := session.Save(); err != nil {
response.RespondFailure(c, "无法更新会话信息,请重试")
return
}
cleanUser, err := setLoginToken(user)
if err != nil {
response.RespondFailure(c, "无法保存会话信息,请重试")
return
}
response.RespondSuccess(c, service.OAuthCallbackResult{Status: "linked", User: cleanUser})
}
func ListExternalAccounts(c *gin.Context) {
userID := c.GetInt("id")
accounts, err := model.ListExternalAccountsByUserID(userID)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, accounts)
}
func DeleteExternalAccount(c *gin.Context) {
rawID := strings.TrimSpace(c.Param("id"))
parsedID, err := strconv.ParseUint(rawID, 10, 64)
if err != nil || parsedID == 0 {
response.RespondBadRequest(c, "绑定记录 ID 无效")
return
}
if err := model.DeleteExternalAccountForUser(uint(parsedID), c.GetInt("id")); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
func parseAuthSourceID(c *gin.Context) (uint, error) {
raw := c.Param("source_id")
if raw == "" {
raw = c.Param("id")
}
parsed, err := strconv.ParseUint(raw, 10, 64)
if err != nil || parsed == 0 {
return 0, fmt.Errorf("认证源 ID 无效")
}
return uint(parsed), nil
}
func getAuthSourceFromRoute(c *gin.Context) (*model.AuthSource, error) {
raw := strings.TrimSpace(c.Param("source"))
if raw == "" {
raw = strings.TrimSpace(c.Param("source_id"))
}
if raw == "" {
raw = strings.TrimSpace(c.Param("id"))
}
if raw == "" {
return nil, fmt.Errorf("认证源不能为空")
}
if parsed, err := strconv.ParseUint(raw, 10, 64); err == nil && parsed > 0 {
source, err := model.GetAuthSourceByID(uint(parsed))
if err != nil {
return nil, err
}
return source, nil
}
source, err := model.GetAuthSourceByName(raw)
if err != nil {
return nil, err
}
return source, nil
}
func oauthStateSessionKey(sourceID uint) string {
return fmt.Sprintf("oauth_state_%d", sourceID)
}
func oauthFrontendCallbackURL(c *gin.Context, sourceID uint) string {
base := strings.TrimRight(common.ServerAddress, "/")
if base == "" {
scheme := "http"
if c.Request.TLS != nil || c.GetHeader("X-Forwarded-Proto") == "https" {
scheme = "https"
}
host := c.Request.Host
if forwardedHost := c.GetHeader("X-Forwarded-Host"); forwardedHost != "" {
host = forwardedHost
}
base = scheme + "://" + host
}
source, err := model.GetAuthSourceByID(sourceID)
sourceName := strconv.FormatUint(uint64(sourceID), 10)
if err == nil && strings.TrimSpace(source.Name) != "" {
sourceName = source.Name
}
callback, _ := url.JoinPath(base, "oauth", sourceName)
return callback
}
@@ -0,0 +1,49 @@
package bind
import (
"encoding/json"
"errors"
"io"
"strconv"
"github.com/gin-gonic/gin"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
)
// DecodeJSONBody decodes JSON reader to target
func DecodeJSONBody(body io.Reader, target any) error {
return json.NewDecoder(body).Decode(target)
}
// OptionalJSON decodes optional JSON body of reader to target, allowing EOF
func OptionalJSON(body io.Reader, target any) error {
if err := json.NewDecoder(body).Decode(target); err != nil && !errors.Is(err, io.EOF) {
return err
}
return nil
}
// IDParam parses "id" parameter from context path
func IDParam(c *gin.Context) (uint, bool) {
return IDParamByName(c, "id")
}
// IDParamByName parses target parameter from context path
func IDParamByName(c *gin.Context, name string) (uint, bool) {
id, err := strconv.ParseUint(c.Param(name), 10, 64)
if err != nil || id == 0 {
response.RespondBadRequest(c, "")
return 0, false
}
return uint(id), true
}
// JSON binds JSON body of context request to target
func JSON(c *gin.Context, target any) bool {
if err := DecodeJSONBody(c.Request.Body, target); err != nil {
response.RespondBadRequest(c, "")
return false
}
return true
}
@@ -0,0 +1,51 @@
package controller
import (
"net/http"
"github.com/gin-gonic/gin"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/rain-kl/openflare/openflare-server/internal/utils/cap"
)
// GetCapChallenge generates a new CAPTCHA challenge
func GetCapChallenge(c *gin.Context) {
scope := c.Param("scope")
if scope == "" {
scope = c.Query("scope")
}
resp, err := service.CapManager.Generate(scope)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{
"success": false,
"error": err.Error(),
})
return
}
c.JSON(http.StatusOK, resp)
}
// RedeemCapChallenge validates CAPTCHA solutions and yields a one-time redeem token
func RedeemCapChallenge(c *gin.Context) {
scope := c.Param("scope")
if scope == "" {
scope = c.Query("scope")
}
var req cap.RedeemRequest
if !bind.JSON(c, &req) {
return
}
resp, err := service.CapManager.Redeem(c.Request.Context(), req.Token, req.Solutions, scope)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{
"success": false,
"error": err.Error(),
})
return
}
c.JSON(http.StatusOK, resp)
}
@@ -0,0 +1,165 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
)
// GetConfigVersions godoc
// @Summary List config versions
// @Tags ConfigVersions
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/config-versions/ [get]
func GetConfigVersions(c *gin.Context) {
versions, err := service.ListConfigVersions()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, versions)
}
// GetConfigVersion godoc
// @Summary Get config version detail
// @Tags ConfigVersions
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Version ID"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/config-versions/{id} [get]
func GetConfigVersion(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
version, err := service.GetConfigVersionDetail(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, version)
}
// GetActiveConfigVersion godoc
// @Summary Get active config version
// @Tags ConfigVersions
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/config-versions/active [get]
func GetActiveConfigVersion(c *gin.Context) {
version, err := service.GetActiveConfigVersion()
if err != nil {
response.RespondFailure(c, "当前没有激活版本")
return
}
response.RespondSuccess(c, version)
}
// PreviewConfigVersion godoc
// @Summary Preview config rendering
// @Tags ConfigVersions
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/config-versions/preview [get]
func PreviewConfigVersion(c *gin.Context) {
preview, err := service.PreviewConfigVersion()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, preview)
}
// DiffConfigVersion godoc
// @Summary Diff current draft against active version
// @Tags ConfigVersions
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/config-versions/diff [get]
func DiffConfigVersion(c *gin.Context) {
diff, err := service.DiffConfigVersion()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, diff)
}
// PublishConfigVersion godoc
// @Summary Publish a new config version
// @Tags ConfigVersions
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/config-versions/publish [post]
func PublishConfigVersion(c *gin.Context) {
username := c.GetString("username")
force := c.Query("force") == "true"
result, err := service.PublishConfigVersion(username, force)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, result.Version)
}
// ActivateConfigVersion godoc
// @Summary Activate an existing config version
// @Tags ConfigVersions
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Version ID"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/config-versions/{id}/activate [post]
func ActivateConfigVersion(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
version, err := service.ActivateConfigVersion(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, version)
}
type CleanupConfigVersionRequest struct {
KeepCount int `json:"keep_count" binding:"required,min=3"`
}
// CleanupConfigVersions godoc
// @Summary Cleanup old config versions
// @Tags ConfigVersions
// @Produce json
// @Security OpenFlareTokenAuth
// @Param request body CleanupConfigVersionRequest true "Cleanup request"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/config-versions/cleanup [post]
func CleanupConfigVersions(c *gin.Context) {
var req CleanupConfigVersionRequest
if !bind.JSON(c, &req) {
return
}
deletedCount, err := service.CleanupConfigVersions(req.KeepCount)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessWithExtras(c, map[string]interface{}{"deleted_count": deletedCount}, gin.H{
"message": "清理成功",
})
}
@@ -0,0 +1,173 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
)
type dashboardOverviewPayload struct {
GeneratedAt any `json:"generated_at"`
Summary service.DashboardSummary `json:"summary"`
Traffic service.DashboardTraffic `json:"traffic"`
Capacity service.DashboardCapacity `json:"capacity"`
Distributions dashboardDistributionsPayload `json:"distributions"`
Trends dashboardTrendsPayload `json:"trends"`
Nodes [][]any `json:"nodes"`
}
type dashboardDistributionsPayload struct {
StatusCodes [][]any `json:"status_codes"`
TopDomains [][]any `json:"top_domains"`
SourceCountries [][]any `json:"source_countries"`
}
type dashboardTrendsPayload struct {
Traffic24h [][]any `json:"traffic_24h"`
Capacity24h [][]any `json:"capacity_24h"`
Network24h [][]any `json:"network_24h"`
DiskIO24h [][]any `json:"disk_io_24h"`
}
// GetDashboardOverview godoc
// @Summary Get dashboard overview
// @Tags Dashboard
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/dashboard/overview [get]
func GetDashboardOverview(c *gin.Context) {
view, err := service.GetDashboardOverview()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, compressDashboardOverview(view))
}
func compressDashboardOverview(view *service.DashboardOverviewView) *dashboardOverviewPayload {
if view == nil {
return &dashboardOverviewPayload{
Distributions: dashboardDistributionsPayload{
StatusCodes: [][]any{},
TopDomains: [][]any{},
SourceCountries: [][]any{},
},
Trends: dashboardTrendsPayload{
Traffic24h: [][]any{},
Capacity24h: [][]any{},
Network24h: [][]any{},
DiskIO24h: [][]any{},
},
Nodes: [][]any{},
}
}
return &dashboardOverviewPayload{
GeneratedAt: view.GeneratedAt,
Summary: view.Summary,
Traffic: view.Traffic,
Capacity: view.Capacity,
Distributions: dashboardDistributionsPayload{
StatusCodes: compressDistributionItems(view.Distributions.StatusCodes),
TopDomains: compressDistributionItems(view.Distributions.TopDomains),
SourceCountries: compressDistributionItems(view.Distributions.SourceCountries),
},
Trends: dashboardTrendsPayload{
Traffic24h: compressTrafficTrendPoints(view.Trends.Traffic24h),
Capacity24h: compressCapacityTrendPoints(view.Trends.Capacity24h),
Network24h: compressNetworkTrendPoints(view.Trends.Network24h),
DiskIO24h: compressDiskIOTrendPoints(view.Trends.DiskIO24h),
},
Nodes: compressDashboardNodes(view.Nodes),
}
}
func compressDistributionItems(items []service.DistributionItem) [][]any {
rows := make([][]any, 0, len(items))
for _, item := range items {
rows = append(rows, []any{item.Key, item.Value})
}
return rows
}
func compressTrafficTrendPoints(points []service.TrafficTrendPoint) [][]any {
rows := make([][]any, 0, len(points))
for _, point := range points {
rows = append(rows, []any{
point.BucketStartedAt,
point.RequestCount,
point.ErrorCount,
point.UniqueVisitorCount,
})
}
return rows
}
func compressCapacityTrendPoints(points []service.CapacityTrendPoint) [][]any {
rows := make([][]any, 0, len(points))
for _, point := range points {
rows = append(rows, []any{
point.BucketStartedAt,
point.AverageCPUUsagePercent,
point.AverageMemoryUsagePercent,
point.ReportedNodes,
})
}
return rows
}
func compressNetworkTrendPoints(points []service.NetworkTrendPoint) [][]any {
rows := make([][]any, 0, len(points))
for _, point := range points {
rows = append(rows, []any{
point.BucketStartedAt,
point.NetworkRxBytes,
point.NetworkTxBytes,
point.OpenrestyRxBytes,
point.OpenrestyTxBytes,
point.ReportedNodes,
})
}
return rows
}
func compressDiskIOTrendPoints(points []service.DiskIOTrendPoint) [][]any {
rows := make([][]any, 0, len(points))
for _, point := range points {
rows = append(rows, []any{
point.BucketStartedAt,
point.DiskReadBytes,
point.DiskWriteBytes,
point.ReportedNodes,
})
}
return rows
}
func compressDashboardNodes(nodes []service.DashboardNodeHealth) [][]any {
rows := make([][]any, 0, len(nodes))
for _, node := range nodes {
rows = append(rows, []any{
node.ID,
node.NodeID,
node.Name,
node.GeoName,
node.GeoLatitude,
node.GeoLongitude,
node.Status,
node.OpenrestyStatus,
node.CurrentVersion,
node.LastSeenAt,
node.ActiveEventCount,
node.CPUUsagePercent,
node.MemoryUsagePercent,
node.StorageUsagePercent,
node.RequestCount,
node.ErrorCount,
node.UniqueVisitorCount,
})
}
return rows
}
@@ -0,0 +1,31 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
)
// CleanupDatabaseObservability godoc
// @Summary Cleanup observability tables
// @Tags Options
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/option/database/cleanup [post]
func CleanupDatabaseObservability(c *gin.Context) {
var input service.DatabaseCleanupInput
if err := bind.OptionalJSON(c.Request.Body, &input); err != nil {
response.RespondBadRequest(c, "")
return
}
result, err := service.CleanupDatabaseObservability(input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, result)
}
@@ -0,0 +1,135 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/gin-gonic/gin"
)
type DnsAccountInput struct {
Name string `json:"name"`
Type string `json:"type"`
Authorization string `json:"authorization"`
}
// GetDnsAccounts godoc
// @Summary List DNS accounts
// @Tags DnsAccounts
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/dns-accounts/ [get]
func GetDnsAccounts(c *gin.Context) {
accounts, err := model.ListDnsAccounts()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, accounts)
}
// CreateDnsAccount godoc
// @Summary Create DNS account
// @Tags DnsAccounts
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Param payload body DnsAccountInput true "DNS account payload"
// @Success 200 {object} map[string]interface{}
// @Router /api/dns-accounts/ [post]
func CreateDnsAccount(c *gin.Context) {
var input DnsAccountInput
if !bind.JSON(c, &input) {
return
}
account := &model.DnsAccount{
Name: input.Name,
Type: input.Type,
Authorization: input.Authorization,
}
if err := account.Insert(); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, account)
}
// UpdateDnsAccount godoc
// @Summary Update DNS account
// @Tags DnsAccounts
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "DNS Account ID"
// @Param payload body DnsAccountInput true "DNS account payload"
// @Success 200 {object} map[string]interface{}
// @Router /api/dns-accounts/{id}/update [post]
func UpdateDnsAccount(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
var input DnsAccountInput
if !bind.JSON(c, &input) {
return
}
account, err := model.GetDnsAccountByID(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
account.Name = input.Name
account.Type = input.Type
account.Authorization = input.Authorization
if err := account.Update(); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, account)
}
// DeleteDnsAccount godoc
// @Summary Delete DNS account
// @Tags DnsAccounts
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "DNS Account ID"
// @Success 200 {object} map[string]interface{}
// @Router /api/dns-accounts/{id}/delete [post]
func DeleteDnsAccount(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
account, err := model.GetDnsAccountByID(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
// Verify no cert uses this before deleting
var count int64
model.DB.Model(&model.TLSCertificate{}).Where("dns_account_id = ?", id).Count(&count)
if count > 0 {
response.RespondFailure(c, "该 DNS 账号已被证书使用,无法删除")
return
}
if err := account.Delete(); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, nil)
}
@@ -0,0 +1,176 @@
package controller
import (
"log/slog"
"net"
"time"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
"golang.org/x/net/websocket"
)
// FlaredHeartbeat godoc
// @Summary Report OpenFlared heartbeat
// @Tags Flared
// @Accept json
// @Produce json
// @Security TunnelTokenAuth
// @Param payload body service.FlaredHeartbeatPayload true "Flared heartbeat payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/flared/heartbeat [post]
func FlaredHeartbeat(c *gin.Context) {
var payload service.FlaredHeartbeatPayload
if !bind.JSON(c, &payload) {
return
}
authNode, ok := c.Get("flared_node")
if !ok {
response.RespondUnauthorized(c, "无权进行此操作,Tunnel Token 无效")
return
}
node := authNode.(*model.Node)
res, err := service.HeartbeatFlared(node, payload)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, res)
}
// FlaredGetActiveConfig godoc
// @Summary Get active tunnel config for OpenFlared
// @Tags Flared
// @Produce json
// @Security TunnelTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/flared/config/active [get]
func FlaredGetActiveConfig(c *gin.Context) {
authNode, ok := c.Get("flared_node")
if !ok {
response.RespondUnauthorized(c, "无权进行此操作,Tunnel Token 无效")
return
}
node := authNode.(*model.Node)
config, err := service.GetFlaredTunnelConfig(node)
if err != nil {
response.RespondFailure(c, "无法生成隧道配置: "+err.Error())
return
}
response.RespondSuccess(c, config)
}
// FlaredReportApplyLog godoc
// @Summary Report OpenFlared apply result
// @Tags Flared
// @Accept json
// @Produce json
// @Security TunnelTokenAuth
// @Param payload body service.ApplyLogPayload true "Apply log payload"
// @Success 200 {object} map[string]interface{}
// @Router /api/flared/apply-log [post]
func FlaredReportApplyLog(c *gin.Context) {
var payload service.ApplyLogPayload
if !bind.JSON(c, &payload) {
return
}
if authNode, ok := c.Get("flared_node"); ok {
payload.NodeID = authNode.(*model.Node).NodeID
}
log, err := service.ReportApplyLog(payload)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, log)
}
// FlaredWebSocket godoc
// @Summary Upgrade OpenFlared connection to websocket
// @Tags Flared
// @Security TunnelTokenAuth
// @Router /api/flared/ws [get]
func FlaredWebSocket(c *gin.Context) {
authNode, ok := c.Get("flared_node")
if !ok {
response.RespondUnauthorized(c, "无权进行此操作,Tunnel Token 无效")
return
}
node := authNode.(*model.Node)
slog.Debug("flared ws upgrade requested", "node_id", node.NodeID, "remote", c.Request.RemoteAddr)
websocket.Handler(func(conn *websocket.Conn) {
client := service.RegisterFlaredWSClient(node.NodeID)
defer service.UnregisterFlaredWSClient(client)
defer func() {
_ = conn.Close()
slog.Debug("flared ws connection closed", "node_id", node.NodeID)
}()
slog.Debug("flared ws upgrade succeeded", "node_id", node.NodeID, "remote", c.Request.RemoteAddr)
go func() {
<-client.Done()
_ = conn.Close()
}()
go streamFlaredWSMessages(c, conn, client)
for {
var message service.WSMessage
_ = conn.SetReadDeadline(time.Now().Add(flaredWSReadTimeout()))
if err := websocket.JSON.Receive(conn, &message); err != nil {
if netErr, ok := err.(net.Error); ok && netErr.Timeout() {
slog.Debug("flared ws receive timeout", "node_id", node.NodeID)
return
}
slog.Debug("flared ws receive failed", "node_id", node.NodeID, "error", err)
return
}
slog.Debug("flared ws message received", "node_id", node.NodeID, "type", message.Type)
switch message.Type {
case "ping":
if !service.SendFlaredWSPong(node.NodeID) {
slog.Debug("flared ws pong enqueue failed", "node_id", node.NodeID)
}
case "pong":
slog.Debug("flared ws pong received", "node_id", node.NodeID)
default:
slog.Debug("flared ws unsupported message type", "node_id", node.NodeID, "type", message.Type)
}
}
}).ServeHTTP(c.Writer, c.Request)
}
func streamFlaredWSMessages(c *gin.Context, conn *websocket.Conn, client *service.WSClient) {
for {
select {
case <-c.Request.Context().Done():
return
case <-client.Done():
return
case message, ok := <-client.Messages():
if !ok {
return
}
_ = conn.SetWriteDeadline(time.Now().Add(agentWSWriteTimeout()))
if err := websocket.JSON.Send(conn, message); err != nil {
slog.Debug("flared ws send failed", "node_id", client.ID(), "error", err)
return
}
}
}
}
func flaredWSReadTimeout() time.Duration {
timeout := time.Duration(common.AgentHeartbeatInterval) * time.Millisecond * 3
if timeout < 30*time.Second {
return 30 * time.Second
}
return timeout
}
@@ -0,0 +1,37 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
)
type geoIPLookupRequest struct {
Provider string `json:"provider"`
IP string `json:"ip"`
}
// LookupGeoIP godoc
// @Summary Test GeoIP lookup
// @Tags Options
// @Accept json
// @Produce json
// @Param payload body geoIPLookupRequest true "GeoIP lookup payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/option/geoip/lookup [post]
func LookupGeoIP(c *gin.Context) {
var request geoIPLookupRequest
if !bind.JSON(c, &request) {
return
}
view, err := service.LookupGeoIP(request.Provider, request.IP)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, view)
}
@@ -0,0 +1,155 @@
package controller
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"log/slog"
"net/http"
"time"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/gin-gonic/gin"
)
type GitHubOAuthResponse struct {
AccessToken string `json:"access_token"`
Scope string `json:"scope"`
TokenType string `json:"token_type"`
}
type GitHubUser struct {
Login string `json:"login"`
Name string `json:"name"`
Email string `json:"email"`
}
func getGitHubUserInfoByCode(code string) (*GitHubUser, error) {
if code == "" {
return nil, errors.New("无效的参数")
}
values := map[string]string{"client_id": common.GitHubClientId, "client_secret": common.GitHubClientSecret, "code": code}
jsonData, err := json.Marshal(values)
if err != nil {
return nil, err
}
req, err := http.NewRequest("POST", "https://github.com/login/oauth/access_token", bytes.NewBuffer(jsonData))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
client := http.Client{
Timeout: 5 * time.Second,
}
res, err := client.Do(req)
if err != nil {
slog.Error("github oauth access token request failed", "error", err)
return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!")
}
defer res.Body.Close()
var oAuthResponse GitHubOAuthResponse
err = json.NewDecoder(res.Body).Decode(&oAuthResponse)
if err != nil {
return nil, err
}
req, err = http.NewRequest("GET", "https://api.github.com/user", nil)
if err != nil {
return nil, err
}
req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", oAuthResponse.AccessToken))
res2, err := client.Do(req)
if err != nil {
slog.Error("github user info request failed", "error", err)
return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!")
}
defer res2.Body.Close()
var githubUser GitHubUser
err = json.NewDecoder(res2.Body).Decode(&githubUser)
if err != nil {
return nil, err
}
if githubUser.Login == "" {
return nil, errors.New("返回值非法,用户字段为空,请稍后重试!")
}
return &githubUser, nil
}
func GitHubOAuth(c *gin.Context) {
if currentUserFromOpenFlareToken(c) != nil {
GitHubBind(c)
return
}
if !common.GitHubOAuthEnabled {
response.RespondFailure(c, "管理员未开启通过 GitHub 登录以及注册")
return
}
code := c.Query("code")
githubUser, err := getGitHubUserInfoByCode(code)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
user := model.User{
GitHubId: githubUser.Login,
}
if model.IsGitHubIdAlreadyTaken(user.GitHubId) {
err := user.FillUserByGitHubId()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
} else {
response.RespondFailure(c, "管理员关闭了新用户注册")
return
}
if user.Status != common.UserStatusEnabled {
response.RespondFailure(c, "用户已被封禁")
return
}
setupLogin(&user, c)
}
func GitHubBind(c *gin.Context) {
if !common.GitHubOAuthEnabled {
response.RespondFailure(c, "管理员未开启通过 GitHub 登录以及注册")
return
}
code := c.Query("code")
githubUser, err := getGitHubUserInfoByCode(code)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
user := model.User{
GitHubId: githubUser.Login,
}
if model.IsGitHubIdAlreadyTaken(user.GitHubId) {
response.RespondFailure(c, "该 GitHub 账户已被绑定")
return
}
currentUser := currentUserFromOpenFlareToken(c)
if currentUser == nil {
response.RespondFailure(c, "无权进行此操作,未登录或 token 无效")
return
}
user.Id = currentUser.Id
err = user.FillUserById()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
user.GitHubId = githubUser.Login
err = user.Update(false)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "bind")
}
@@ -0,0 +1,117 @@
package controller
import (
"strings"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
)
// GetManagedDomains godoc
// @Summary List managed domains
// @Tags ManagedDomains
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/managed-domains/ [get]
func GetManagedDomains(c *gin.Context) {
domains, err := service.ListManagedDomains()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, domains)
}
// CreateManagedDomain godoc
// @Summary Create managed domain
// @Tags ManagedDomains
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Param payload body service.ManagedDomainInput true "Managed domain payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/managed-domains/ [post]
func CreateManagedDomain(c *gin.Context) {
var input service.ManagedDomainInput
if !bind.JSON(c, &input) {
return
}
domain, err := service.CreateManagedDomain(input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, domain)
}
// UpdateManagedDomain godoc
// @Summary Update managed domain
// @Tags ManagedDomains
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Managed domain ID"
// @Param payload body service.ManagedDomainInput true "Managed domain payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/managed-domains/{id}/update [post]
func UpdateManagedDomain(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
var input service.ManagedDomainInput
if !bind.JSON(c, &input) {
return
}
domain, err := service.UpdateManagedDomain(id, input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, domain)
}
// DeleteManagedDomain godoc
// @Summary Delete managed domain
// @Tags ManagedDomains
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Managed domain ID"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/managed-domains/{id}/delete [post]
func DeleteManagedDomain(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
if err := service.DeleteManagedDomain(id); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, nil)
}
// MatchManagedDomainCertificate godoc
// @Summary Match certificate for domain
// @Tags ManagedDomains
// @Produce json
// @Security OpenFlareTokenAuth
// @Param domain query string true "Domain"
// @Success 200 {object} map[string]interface{}
// @Router /api/managed-domains/match [get]
func MatchManagedDomainCertificate(c *gin.Context) {
domain := strings.TrimSpace(c.Query("domain"))
result, err := service.MatchManagedDomainCertificate(domain)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, result)
}
@@ -0,0 +1,148 @@
package controller
import (
"fmt"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/rain-kl/openflare/openflare-server/internal/utils/mail"
"github.com/rain-kl/openflare/openflare-server/internal/utils/security"
"github.com/rain-kl/openflare/openflare-server/internal/utils/validation"
"github.com/gin-gonic/gin"
)
// GetStatus godoc
// @Summary Get server status
// @Tags Public
// @Produce json
// @Success 200 {object} map[string]interface{}
// @Router /api/status [get]
func GetStatus(c *gin.Context) {
authSources, err := service.PublicAuthSources("/api")
if err != nil {
authSources = []service.PublicAuthSource{}
}
response.RespondSuccess(c, gin.H{
"version": common.Version,
"start_time": common.StartTime,
"email_verification": common.EmailVerificationEnabled,
"github_oauth": common.GitHubOAuthEnabled,
"github_client_id": common.GitHubClientId,
"system_name": common.SystemName,
"home_page_link": common.HomePageLink,
"footer_html": common.Footer,
"wechat_qrcode": common.WeChatAccountQRCodeImageURL,
"wechat_login": common.WeChatAuthEnabled,
"server_address": common.ServerAddress,
"password_register_enabled": common.PasswordRegisterEnabled,
"cap_login_enabled": common.CapLoginEnabled,
"auth_sources": authSources,
})
}
func GetNotice(c *gin.Context) {
common.OptionMapRWMutex.RLock()
defer common.OptionMapRWMutex.RUnlock()
response.RespondSuccess(c, common.OptionMap["Notice"])
}
func GetAbout(c *gin.Context) {
common.OptionMapRWMutex.RLock()
defer common.OptionMapRWMutex.RUnlock()
response.RespondSuccess(c, common.OptionMap["About"])
}
func SendEmailVerification(c *gin.Context) {
email := c.Query("email")
if err := validation.Validate.Var(email, "required,email"); err != nil {
response.RespondFailure(c, "无效的参数")
return
}
if model.IsEmailAlreadyTaken(email) {
response.RespondFailure(c, "邮箱地址已被占用")
return
}
code := security.GenerateVerificationCode(6)
security.RegisterVerificationCodeWithKey(email, code, security.EmailVerificationPurpose)
subject := fmt.Sprintf("%s邮箱验证邮件", common.SystemName)
content := fmt.Sprintf("<p>您好,你正在进行%s邮箱验证。</p>"+
"<p>您的验证码为: <strong>%s</strong></p>"+
"<p>验证码 %d 分钟内有效,如果不是本人操作,请忽略。</p>", common.SystemName, code, security.VerificationValidMinutes)
cfg := mail.SMTPConfig{
Server: common.SMTPServer,
Port: common.SMTPPort,
Account: common.SMTPAccount,
Token: common.SMTPToken,
SystemName: common.SystemName,
}
err := mail.SendEmail(cfg, subject, email, content)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
func SendPasswordResetEmail(c *gin.Context) {
email := c.Query("email")
if err := validation.Validate.Var(email, "required,email"); err != nil {
response.RespondFailure(c, "无效的参数")
return
}
if !model.IsEmailAlreadyTaken(email) {
response.RespondFailure(c, "该邮箱地址未注册")
return
}
code := security.GenerateVerificationCode(0)
security.RegisterVerificationCodeWithKey(email, code, security.PasswordResetPurpose)
link := fmt.Sprintf("%s/user/reset?email=%s&token=%s", common.ServerAddress, email, code)
subject := fmt.Sprintf("%s密码重置", common.SystemName)
content := fmt.Sprintf("<p>您好,你正在进行%s密码重置。</p>"+
"<p>点击<a href='%s'>此处</a>进行密码重置。</p>"+
"<p>重置链接 %d 分钟内有效,如果不是本人操作,请忽略。</p>", common.SystemName, link, security.VerificationValidMinutes)
cfg := mail.SMTPConfig{
Server: common.SMTPServer,
Port: common.SMTPPort,
Account: common.SMTPAccount,
Token: common.SMTPToken,
SystemName: common.SystemName,
}
err := mail.SendEmail(cfg, subject, email, content)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
type PasswordResetRequest struct {
Email string `json:"email"`
Token string `json:"token"`
}
func ResetPassword(c *gin.Context) {
var req PasswordResetRequest
if !bind.JSON(c, &req) {
return
}
if req.Email == "" || req.Token == "" {
response.RespondFailure(c, "无效的参数")
return
}
if !security.VerifyCodeWithKey(req.Email, req.Token, security.PasswordResetPurpose) {
response.RespondFailure(c, "重置链接非法或已过期")
return
}
password := security.GenerateVerificationCode(12)
err := model.ResetUserPasswordByEmail(req.Email, password)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
security.DeleteKey(req.Email, security.PasswordResetPurpose)
response.RespondSuccess(c, password)
}
@@ -0,0 +1,288 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
)
type nodeAgentUpdateRequest struct {
Channel string `json:"channel"`
TagName string `json:"tag_name"`
}
type nodeObservabilityQuery struct {
Hours int `form:"hours"`
Limit int `form:"limit"`
}
// CreateNode godoc
// @Summary Create node
// @Tags Nodes
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Param payload body service.NodeInput true "Node payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/ [post]
func CreateNode(c *gin.Context) {
var input service.NodeInput
if !bind.JSON(c, &input) {
return
}
node, err := service.CreateNode(input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, node)
}
// GetNodeBootstrapToken godoc
// @Summary Get global discovery token
// @Tags Nodes
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/nodes/bootstrap-token [get]
func GetNodeBootstrapToken(c *gin.Context) {
bootstrap, err := service.GetNodeBootstrapView()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, bootstrap)
}
// RotateNodeBootstrapToken godoc
// @Summary Rotate global discovery token
// @Tags Nodes
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/nodes/bootstrap-token/rotate [post]
func RotateNodeBootstrapToken(c *gin.Context) {
bootstrap, err := service.RotateGlobalDiscoveryToken()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, bootstrap)
}
// UpdateNode godoc
// @Summary Update node
// @Tags Nodes
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Node ID"
// @Param payload body service.NodeInput true "Node payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/update [post]
func UpdateNode(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
var input service.NodeInput
if !bind.JSON(c, &input) {
return
}
node, err := service.UpdateNode(id, input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, node)
}
// DeleteNode godoc
// @Summary Delete node
// @Tags Nodes
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Node ID"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/delete [post]
func DeleteNode(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
if err := service.DeleteNode(id); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
// RequestNodeAgentUpdate godoc
// @Summary Request agent self-update on node
// @Tags Nodes
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Node ID"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/agent-update [post]
func RequestNodeAgentUpdate(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
var request nodeAgentUpdateRequest
if c.Request.ContentLength > 0 {
if err := bind.OptionalJSON(c.Request.Body, &request); err != nil {
response.RespondBadRequest(c, "")
return
}
}
node, err := service.RequestNodeAgentUpdate(id, service.NodeAgentUpdateInput{
Channel: request.Channel,
TagName: request.TagName,
})
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, node)
}
// RequestNodeOpenrestyRestart godoc
// @Summary Request openresty restart on node
// @Tags Nodes
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Node ID"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/openresty-restart [post]
func RequestNodeOpenrestyRestart(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
node, err := service.RequestNodeOpenrestyRestart(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, node)
}
// RequestNodeForceSync godoc
// @Summary Request force sync config on node
// @Tags Nodes
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Node ID"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/force-sync [post]
func RequestNodeForceSync(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
node, err := service.RequestNodeForceSync(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, node)
}
// GetNodeAgentRelease godoc
// @Summary Check latest agent release for node
// @Tags Nodes
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Node ID"
// @Param channel query string false "stable or preview"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/agent-release [get]
func GetNodeAgentRelease(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
release, err := service.GetNodeAgentRelease(c.Request.Context(), id, c.Query("channel"))
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, release)
}
// GetNodeObservability godoc
// @Summary Get node observability details
// @Tags Nodes
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Node ID"
// @Param hours query int false "Lookback window in hours"
// @Param limit query int false "Max records per section"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/observability [get]
func GetNodeObservability(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
var query nodeObservabilityQuery
if err := c.ShouldBindQuery(&query); err != nil {
response.RespondBadRequest(c, "")
return
}
view, err := service.GetNodeObservability(id, service.NodeObservabilityQuery{
Hours: query.Hours,
Limit: query.Limit,
})
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, view)
}
// CleanupNodeHealthEvents godoc
// @Summary Cleanup node health events
// @Tags Nodes
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Node ID"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/observability/cleanup [post]
func CleanupNodeHealthEvents(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
result, err := service.CleanupNodeHealthEvents(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, result)
}
@@ -0,0 +1,417 @@
package controller
import (
"fmt"
"regexp"
"strconv"
"strings"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/rain-kl/openflare/pkg/geoip"
"github.com/rain-kl/openflare/pkg/utils"
"github.com/gin-gonic/gin"
)
var (
openRestySizePattern = regexp.MustCompile(`^\d+[kKmMgG]?$`)
openRestyProxyBuffersPattern = regexp.MustCompile(`^\d+\s+\d+[kKmMgG]?$`)
openRestyCacheLevelsPattern = regexp.MustCompile(`^\d{1,2}(?::\d{1,2}){0,2}$`)
openRestyDurationTokenPattern = regexp.MustCompile(`^\d+[smhdwSMHDW]$`)
)
type optionBatchPayload struct {
Options []model.Option `json:"options"`
}
func validateRateLimitOption(key string, value string) error {
maxDurationSeconds := int(common.RateLimitKeyExpirationDuration.Seconds())
switch key {
case "GlobalApiRateLimitNum", "GlobalWebRateLimitNum", "CriticalRateLimitNum":
intValue, err := strconv.Atoi(value)
if err != nil || intValue <= 0 {
return fmt.Errorf("%s 必须为大于 0 的整数", key)
}
return nil
case "GlobalApiRateLimitDuration", "GlobalWebRateLimitDuration", "CriticalRateLimitDuration":
intValue, err := strconv.Atoi(value)
if err != nil || intValue <= 0 {
return fmt.Errorf("%s 必须为大于 0 的整数秒", key)
}
if intValue > maxDurationSeconds {
return fmt.Errorf("%s 不能大于 %d 秒", key, maxDurationSeconds)
}
return nil
default:
return nil
}
}
func validatePositiveIntegerOption(key string, value string) error {
intValue, err := strconv.Atoi(value)
if err != nil || intValue <= 0 {
return fmt.Errorf("%s 必须为大于 0 的整数", key)
}
return nil
}
func validateBooleanOption(key string, value string) error {
switch value {
case "true", "false":
return nil
default:
return fmt.Errorf("%s 必须为 true 或 false", key)
}
}
func validateGeoIPOption(key string, value string) error {
if key != "GeoIPProvider" {
return nil
}
if !geoip.IsValidProvider(value) {
return fmt.Errorf("%s 仅支持 disabled、mmdb、ip-api、geojs、ipinfo", key)
}
return nil
}
func validateDatabaseCleanupOption(key string, value string) error {
switch key {
case "DatabaseAutoCleanupEnabled":
return validateBooleanOption(key, value)
case "DatabaseAutoCleanupRetentionDays":
intValue, err := strconv.Atoi(value)
if err != nil || intValue < 1 {
return fmt.Errorf("%s 必须为大于等于 1 的整数天", key)
}
return nil
default:
return nil
}
}
func validateAgentOption(key string, value string) error {
switch key {
case "AgentWebsocketUpgradeEnabled":
return validateBooleanOption(key, strings.TrimSpace(value))
default:
return nil
}
}
func validateUptimeKumaOption(key string, value string, state map[string]string) error {
trimmed := strings.TrimSpace(value)
switch key {
case "UptimeKumaEnabled":
if err := validateBooleanOption(key, trimmed); err != nil {
return err
}
if trimmed == "true" {
url := strings.TrimSpace(state["UptimeKumaUrl"])
username := strings.TrimSpace(state["UptimeKumaUsername"])
password := strings.TrimSpace(state["UptimeKumaPassword"])
if url == "" {
return fmt.Errorf("启用 Uptime Kuma 时地址不能为空")
}
if username == "" {
return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空")
}
if password == "" && common.UptimeKumaPassword == "" {
return fmt.Errorf("启用 Uptime Kuma 时密码不能为空")
}
}
case "UptimeKumaUsername":
if trimmed == "" && state["UptimeKumaEnabled"] == "true" {
return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空")
}
case "UptimeKumaPassword":
// No specific format checks needed
case "UptimeKumaUrl":
if trimmed != "" {
if !strings.HasPrefix(trimmed, "http://") && !strings.HasPrefix(trimmed, "https://") {
return fmt.Errorf("Uptime Kuma 地址必须以 http:// 或 https:// 开头")
}
}
case "UptimeKumaMonitorScope":
if trimmed != "all" && trimmed != "selected" {
return fmt.Errorf("监控范围必须为全部站点 (all) 或选择站点 (selected)")
}
case "UptimeKumaSyncInterval", "UptimeKumaInterval", "UptimeKumaRetryInterval", "UptimeKumaTimeout":
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
return err
}
case "UptimeKumaRetry":
intValue, err := strconv.Atoi(trimmed)
if err != nil || intValue < 0 {
return fmt.Errorf("%s 必须为大于等于 0 的整数", key)
}
}
return nil
}
func validateOpenRestyOption(key string, value string) error {
trimmed := strings.TrimSpace(value)
switch key {
case "OpenRestyDefaultServerReturnStatus":
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
return err
}
statusCode, _ := strconv.Atoi(trimmed)
if statusCode < 100 || statusCode > 999 {
return fmt.Errorf("%s 必须在 100 到 999 之间", key)
}
return nil
case "OpenRestyWorkerProcesses":
if trimmed == "auto" {
return nil
}
return validatePositiveIntegerOption(key, trimmed)
case "OpenRestyWorkerConnections",
"OpenRestyWorkerRlimitNofile",
"OpenRestyKeepaliveTimeout",
"OpenRestyKeepaliveRequests",
"OpenRestyClientHeaderTimeout",
"OpenRestyClientBodyTimeout",
"OpenRestySendTimeout",
"OpenRestyProxyConnectTimeout",
"OpenRestyProxySendTimeout",
"OpenRestyProxyReadTimeout",
"OpenRestyGzipMinLength":
return validatePositiveIntegerOption(key, trimmed)
case "OpenRestyGzipCompLevel":
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
return err
}
level, _ := strconv.Atoi(trimmed)
if level > 9 {
return fmt.Errorf("%s 不能大于 9", key)
}
return nil
case "OpenRestyEventsUse":
if trimmed == "" {
return nil
}
switch trimmed {
case "epoll", "kqueue", "poll", "select", "rtsig", "/dev/poll", "eventport":
return nil
default:
return fmt.Errorf("%s 仅支持 epoll、kqueue、poll、select、rtsig、/dev/poll、eventport 或留空", key)
}
case "OpenRestyResolvers":
if trimmed == "" {
return nil
}
if !regexp.MustCompile(`^[a-zA-Z0-9.:\-\s]+$`).MatchString(trimmed) {
return fmt.Errorf("%s 包含非法字符,请填入有效的 IP 地址或域名,以空格分隔", key)
}
return nil
case "OpenRestyEventsMultiAcceptEnabled",
"OpenRestyWebsocketEnabled",
"OpenRestyHTTP3Enabled",
"OpenRestyProxyRequestBufferingEnabled",
"OpenRestyProxyBufferingEnabled",
"OpenRestyGzipEnabled",
"OpenRestyCacheEnabled",
"OpenRestyCacheLockEnabled":
return validateBooleanOption(key, trimmed)
case "OpenRestyProxyBuffers", "OpenRestyLargeClientHeaderBuffers":
if openRestyProxyBuffersPattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须类似 \"16 16k\"", key)
case "OpenRestyProxyBufferSize", "OpenRestyProxyBusyBuffersSize", "OpenRestyCacheMaxSize", "OpenRestyClientMaxBodySize":
if openRestySizePattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须为整数或带 k/m/g 单位的大小值", key)
case "OpenRestyCachePath":
if strings.ContainsAny(trimmed, "\r\n\t") {
return fmt.Errorf("%s 不能包含换行或制表符", key)
}
return nil
case "OpenRestyCacheLevels":
if openRestyCacheLevelsPattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须类似 \"1:2\" 或 \"1:2:2\"", key)
case "OpenRestyCacheInactive", "OpenRestyCacheLockTimeout":
if openRestyDurationTokenPattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须为带单位的时长,例如 30m 或 5s", key)
case "OpenRestyCacheKeyTemplate":
if trimmed == "" {
return fmt.Errorf("%s 不能为空", key)
}
if strings.ContainsAny(trimmed, "\r\n") {
return fmt.Errorf("%s 不能包含换行", key)
}
return nil
case "OpenRestyCacheUseStale":
if trimmed == "" {
return fmt.Errorf("%s 不能为空", key)
}
allowedTokens := map[string]struct{}{
"error": {}, "timeout": {}, "invalid_header": {}, "updating": {},
"http_500": {}, "http_502": {}, "http_503": {}, "http_504": {},
"http_403": {}, "http_404": {}, "http_429": {}, "off": {},
}
for _, token := range strings.Fields(trimmed) {
if _, ok := allowedTokens[token]; !ok {
return fmt.Errorf("%s 包含不支持的值 %q", key, token)
}
}
return nil
case "OpenRestyMainConfigTemplate":
return service.ValidateOpenRestyMainConfigTemplate(value)
default:
return nil
}
}
func buildOptionValidationState(options []model.Option) map[string]string {
common.OptionMapRWMutex.RLock()
state := make(map[string]string, len(common.OptionMap)+len(options))
for key, value := range common.OptionMap {
state[key] = value
}
common.OptionMapRWMutex.RUnlock()
for _, option := range options {
state[option.Key] = option.Value
}
return state
}
func validateOptionWithState(option model.Option, state map[string]string) error {
switch option.Key {
case "GitHubOAuthEnabled":
if option.Value == "true" && strings.TrimSpace(state["GitHubClientId"]) == "" {
return fmt.Errorf("无法启用 GitHub OAuth,请先填入 GitHub Client ID 以及 GitHub Client Secret!")
}
case "WeChatAuthEnabled":
if option.Value == "true" && strings.TrimSpace(state["WeChatServerAddress"]) == "" {
return fmt.Errorf("无法启用微信登录,请先填入微信登录相关配置信息!")
}
}
if err := validateRateLimitOption(option.Key, option.Value); err != nil {
return err
}
if err := validateOpenRestyOption(option.Key, option.Value); err != nil {
return err
}
if err := validateGeoIPOption(option.Key, option.Value); err != nil {
return err
}
if err := validateDatabaseCleanupOption(option.Key, option.Value); err != nil {
return err
}
if err := validateAgentOption(option.Key, option.Value); err != nil {
return err
}
if err := validateUptimeKumaOption(option.Key, option.Value, state); err != nil {
return err
}
return nil
}
func updateOptions(options []model.Option) error {
if len(options) == 0 {
return fmt.Errorf("无效的参数")
}
state := buildOptionValidationState(options)
for _, option := range options {
if strings.TrimSpace(option.Key) == "" {
return fmt.Errorf("无效的参数")
}
if err := validateOptionWithState(option, state); err != nil {
return err
}
}
return model.UpdateOptions(options)
}
// GetOptions godoc
// @Summary List editable options
// @Tags Options
// @Produce json
// @Success 200 {object} map[string]interface{}
// @Router /api/option/ [get]
func GetOptions(c *gin.Context) {
var options []*model.Option
common.OptionMapRWMutex.RLock()
for k, v := range common.OptionMap {
if strings.Contains(k, "Token") || strings.Contains(k, "Secret") || strings.Contains(k, "Password") {
continue
}
options = append(options, &model.Option{
Key: k,
Value: utils.Interface2String(v),
})
}
common.OptionMapRWMutex.RUnlock()
response.RespondSuccess(c, options)
}
// UpdateOption godoc
// @Summary Update option
// @Tags Options
// @Accept json
// @Produce json
// @Param payload body model.Option true "Option payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/option/update [post]
func UpdateOption(c *gin.Context) {
var option model.Option
if !bind.JSON(c, &option) {
return
}
state := buildOptionValidationState([]model.Option{option})
if err := validateOptionWithState(option, state); err != nil {
response.RespondFailure(c, err.Error())
return
}
err := model.UpdateOption(option.Key, option.Value)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
// UpdateOptionsBatch godoc
// @Summary Batch update options
// @Tags Options
// @Accept json
// @Produce json
// @Param payload body optionBatchPayload true "Batch option payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/option/update-batch [post]
func UpdateOptionsBatch(c *gin.Context) {
var payload optionBatchPayload
if !bind.JSON(c, &payload) {
return
}
if len(payload.Options) == 0 {
response.RespondBadRequest(c, "无效的参数")
return
}
if err := updateOptions(payload.Options); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
@@ -0,0 +1,119 @@
package controller
import (
"testing"
)
func TestValidateOpenRestyOption(t *testing.T) {
testCases := []struct {
name string
key string
value string
wantErr bool
}{
{name: "default server status valid 421", key: "OpenRestyDefaultServerReturnStatus", value: "421"},
{name: "default server status valid 200", key: "OpenRestyDefaultServerReturnStatus", value: "200"},
{name: "default server status invalid 99", key: "OpenRestyDefaultServerReturnStatus", value: "99", wantErr: true},
{name: "default server status invalid 1000", key: "OpenRestyDefaultServerReturnStatus", value: "1000", wantErr: true},
{name: "default server status invalid abc", key: "OpenRestyDefaultServerReturnStatus", value: "abc", wantErr: true},
{name: "worker processes auto", key: "OpenRestyWorkerProcesses", value: "auto"},
{name: "worker processes number", key: "OpenRestyWorkerProcesses", value: "8"},
{name: "worker processes invalid", key: "OpenRestyWorkerProcesses", value: "0", wantErr: true},
{name: "events use empty", key: "OpenRestyEventsUse", value: ""},
{name: "events use invalid", key: "OpenRestyEventsUse", value: "io_uring", wantErr: true},
{name: "resolvers valid", key: "OpenRestyResolvers", value: "1.1.1.1 8.8.8.8"},
{name: "resolvers invalid", key: "OpenRestyResolvers", value: "1.1.1.1; 8.8.8.8", wantErr: true},
{name: "proxy buffers valid", key: "OpenRestyProxyBuffers", value: "16 16k"},
{name: "proxy buffers invalid", key: "OpenRestyProxyBuffers", value: "16x16k", wantErr: true},
{name: "cache max size valid", key: "OpenRestyCacheMaxSize", value: "2g"},
{name: "cache max size invalid", key: "OpenRestyCacheMaxSize", value: "2gb", wantErr: true},
{name: "client max body size valid", key: "OpenRestyClientMaxBodySize", value: "64m"},
{name: "client max body size invalid", key: "OpenRestyClientMaxBodySize", value: "64mb", wantErr: true},
{name: "large client header buffers valid", key: "OpenRestyLargeClientHeaderBuffers", value: "4 16k"},
{name: "large client header buffers invalid", key: "OpenRestyLargeClientHeaderBuffers", value: "4x16k", wantErr: true},
{name: "proxy request buffering valid", key: "OpenRestyProxyRequestBufferingEnabled", value: "true"},
{name: "proxy request buffering invalid", key: "OpenRestyProxyRequestBufferingEnabled", value: "on", wantErr: true},
{name: "websocket valid", key: "OpenRestyWebsocketEnabled", value: "false"},
{name: "websocket invalid", key: "OpenRestyWebsocketEnabled", value: "off", wantErr: true},
{name: "cache inactive valid", key: "OpenRestyCacheInactive", value: "30m"},
{name: "cache inactive invalid", key: "OpenRestyCacheInactive", value: "30", wantErr: true},
{name: "cache use stale valid", key: "OpenRestyCacheUseStale", value: "error timeout http_500"},
{name: "cache use stale invalid", key: "OpenRestyCacheUseStale", value: "error whatever", wantErr: true},
{name: "gzip level valid", key: "OpenRestyGzipCompLevel", value: "9"},
{name: "gzip level invalid", key: "OpenRestyGzipCompLevel", value: "10", wantErr: true},
}
for _, testCase := range testCases {
err := validateOpenRestyOption(testCase.key, testCase.value)
if testCase.wantErr && err == nil {
t.Fatalf("%s: expected error", testCase.name)
}
if !testCase.wantErr && err != nil {
t.Fatalf("%s: unexpected error: %v", testCase.name, err)
}
}
}
func TestValidateAgentOption(t *testing.T) {
if err := validateAgentOption("AgentWebsocketUpgradeEnabled", "true"); err != nil {
t.Fatalf("expected websocket upgrade option to accept true: %v", err)
}
if err := validateAgentOption("AgentWebsocketUpgradeEnabled", "false"); err != nil {
t.Fatalf("expected websocket upgrade option to accept false: %v", err)
}
if err := validateAgentOption("AgentWebsocketUpgradeEnabled", "on"); err == nil {
t.Fatal("expected websocket upgrade option to reject non-boolean value")
}
}
func TestValidateUptimeKumaOption(t *testing.T) {
state := map[string]string{
"UptimeKumaUrl": "http://localhost:3001",
"UptimeKumaUsername": "admin",
"UptimeKumaPassword": "password",
}
testCases := []struct {
name string
key string
value string
wantErr bool
}{
{name: "enabled true", key: "UptimeKumaEnabled", value: "true"},
{name: "enabled false", key: "UptimeKumaEnabled", value: "false"},
{name: "enabled invalid", key: "UptimeKumaEnabled", value: "on", wantErr: true},
{name: "url http valid", key: "UptimeKumaUrl", value: "http://192.168.1.100:3001"},
{name: "url https valid", key: "UptimeKumaUrl", value: "https://kuma.example.com"},
{name: "url invalid", key: "UptimeKumaUrl", value: "kuma.example.com", wantErr: true},
{name: "scope all", key: "UptimeKumaMonitorScope", value: "all"},
{name: "scope selected", key: "UptimeKumaMonitorScope", value: "selected"},
{name: "scope invalid", key: "UptimeKumaMonitorScope", value: "none", wantErr: true},
{name: "sync interval valid", key: "UptimeKumaSyncInterval", value: "5"},
{name: "sync interval invalid", key: "UptimeKumaSyncInterval", value: "0", wantErr: true},
{name: "interval valid", key: "UptimeKumaInterval", value: "60"},
{name: "interval invalid", key: "UptimeKumaInterval", value: "-60", wantErr: true},
{name: "retry valid", key: "UptimeKumaRetry", value: "0"},
{name: "retry positive valid", key: "UptimeKumaRetry", value: "3"},
{name: "retry invalid", key: "UptimeKumaRetry", value: "-1", wantErr: true},
}
for _, tc := range testCases {
err := validateUptimeKumaOption(tc.key, tc.value, state)
if tc.wantErr && err == nil {
t.Fatalf("%s: expected error", tc.name)
}
if !tc.wantErr && err != nil {
t.Fatalf("%s: unexpected error: %v", tc.name, err)
}
}
// Test enabling Uptime Kuma when URL or credentials are empty in state
stateEmpty := map[string]string{
"UptimeKumaUrl": "",
"UptimeKumaUsername": "",
"UptimeKumaPassword": "",
}
if err := validateUptimeKumaOption("UptimeKumaEnabled", "true", stateEmpty); err == nil {
t.Fatal("expected error when enabling Uptime Kuma with empty URL/credentials in state")
}
}
@@ -0,0 +1,73 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
)
func GetOrigins(c *gin.Context) {
origins, err := service.ListOrigins()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, origins)
}
func GetOrigin(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
origin, err := service.GetOriginDetail(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, origin)
}
func CreateOrigin(c *gin.Context) {
var input service.OriginInput
if !bind.JSON(c, &input) {
return
}
origin, err := service.CreateOrigin(input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, origin)
}
func UpdateOrigin(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
var input service.OriginInput
if !bind.JSON(c, &input) {
return
}
origin, err := service.UpdateOrigin(id, input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, origin)
}
func DeleteOrigin(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
if err := service.DeleteOrigin(id); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, nil)
}
@@ -0,0 +1,170 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
)
func ListPagesProjects(c *gin.Context) {
projects, err := service.ListPagesProjects()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, projects)
}
func GetPagesProject(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
project, err := service.GetPagesProject(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, project)
}
func CreatePagesProject(c *gin.Context) {
var input service.PagesProjectInput
if !bind.JSON(c, &input) {
return
}
project, err := service.CreatePagesProject(input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, project)
}
func UpdatePagesProject(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
var input service.PagesProjectInput
if !bind.JSON(c, &input) {
return
}
project, err := service.UpdatePagesProject(id, input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, project)
}
func DeletePagesProject(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
if err := service.DeletePagesProject(id); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, nil)
}
func ListPagesDeployments(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
deployments, err := service.ListPagesProjectDeployments(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, deployments)
}
func UploadPagesDeployment(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
file, err := c.FormFile("package")
if err != nil {
response.RespondBadRequest(c, "缺少 Pages 部署包")
return
}
deployment, err := service.UploadPagesDeployment(
id,
file,
c.PostForm("root_dir"),
c.PostForm("entry_file"),
c.GetString("username"),
)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, deployment)
}
func ActivatePagesDeployment(c *gin.Context) {
projectID, ok := bind.IDParam(c)
if !ok {
return
}
deploymentID, ok := bind.IDParamByName(c, "deployment_id")
if !ok {
return
}
project, err := service.ActivatePagesDeployment(projectID, deploymentID)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, project)
}
func DeletePagesDeployment(c *gin.Context) {
projectID, ok := bind.IDParam(c)
if !ok {
return
}
deploymentID, ok := bind.IDParamByName(c, "deployment_id")
if !ok {
return
}
if err := service.DeletePagesDeployment(projectID, deploymentID); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, nil)
}
func ListPagesDeploymentFiles(c *gin.Context) {
deploymentID, ok := bind.IDParamByName(c, "deployment_id")
if !ok {
return
}
files, err := service.ListPagesDeploymentFiles(deploymentID)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, files)
}
func AgentDownloadPagesDeploymentPackage(c *gin.Context) {
deploymentID, ok := bind.IDParamByName(c, "deployment_id")
if !ok {
return
}
filePath, fileName, err := service.GetPagesDeploymentPackagePath(deploymentID)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
c.Header("Content-Disposition", "attachment; filename="+fileName)
c.File(filePath)
}
@@ -0,0 +1,119 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
)
// GetProxyRoutes godoc
// @Summary List proxy routes
// @Tags ProxyRoutes
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/proxy-routes/ [get]
func GetProxyRoutes(c *gin.Context) {
routes, err := service.ListProxyRoutes()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, routes)
}
// GetProxyRoute godoc
// @Summary Get proxy route detail
// @Tags ProxyRoutes
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Route ID"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/proxy-routes/{id} [get]
func GetProxyRoute(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
route, err := service.GetProxyRoute(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, route)
}
// CreateProxyRoute godoc
// @Summary Create proxy route
// @Tags ProxyRoutes
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Param payload body service.ProxyRouteInput true "Proxy route payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/proxy-routes/ [post]
func CreateProxyRoute(c *gin.Context) {
var input service.ProxyRouteInput
if !bind.JSON(c, &input) {
return
}
route, err := service.CreateProxyRoute(input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, route)
}
// UpdateProxyRoute godoc
// @Summary Update proxy route
// @Tags ProxyRoutes
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Route ID"
// @Param payload body service.ProxyRouteInput true "Proxy route payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/proxy-routes/{id}/update [post]
func UpdateProxyRoute(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
var input service.ProxyRouteInput
if !bind.JSON(c, &input) {
return
}
route, err := service.UpdateProxyRoute(id, input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, route)
}
// DeleteProxyRoute godoc
// @Summary Delete proxy route
// @Tags ProxyRoutes
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Route ID"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/proxy-routes/{id}/delete [post]
func DeleteProxyRoute(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
if err := service.DeleteProxyRoute(id); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, nil)
}
@@ -0,0 +1,121 @@
package controller
import (
"log/slog"
"net"
"time"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
"golang.org/x/net/websocket"
)
// RelayHeartbeat godoc
// @Summary Report relay heartbeat
// @Tags Relay
// @Accept json
// @Produce json
// @Security AccessTokenAuth
// @Param payload body service.RelayHeartbeatPayload true "Relay heartbeat payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/relay/heartbeat [post]
func RelayHeartbeat(c *gin.Context) {
var payload service.RelayHeartbeatPayload
if !bind.JSON(c, &payload) {
return
}
payload.IP = service.ResolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
authNode, ok := c.Get("relay_node")
if !ok {
response.RespondUnauthorized(c, "无权进行此操作")
return
}
node := authNode.(*model.Node)
result, err := service.HeartbeatRelay(node, payload)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, result)
}
// RelayWebSocket godoc
// @Summary Upgrade relay connection to websocket
// @Tags Relay
// @Security AccessTokenAuth
// @Router /api/relay/ws [get]
func RelayWebSocket(c *gin.Context) {
authNode, ok := c.Get("relay_node")
if !ok {
response.RespondUnauthorized(c, "无权进行此操作")
return
}
node := authNode.(*model.Node)
slog.Debug("relay ws upgrade requested", "node_id", node.NodeID, "remote", c.Request.RemoteAddr)
websocket.Handler(func(conn *websocket.Conn) {
client := service.RegisterRelayWSClient(node.NodeID)
defer service.UnregisterRelayWSClient(client)
defer func() {
_ = conn.Close()
slog.Debug("relay ws connection closed", "node_id", node.NodeID)
}()
slog.Debug("relay ws upgrade succeeded", "node_id", node.NodeID, "remote", c.Request.RemoteAddr)
go func() {
<-client.Done()
_ = conn.Close()
}()
go streamRelayWSMessages(c, conn, client)
for {
var message service.WSMessage
_ = conn.SetReadDeadline(time.Now().Add(agentWSReadTimeout()))
if err := websocket.JSON.Receive(conn, &message); err != nil {
if netErr, ok := err.(net.Error); ok && netErr.Timeout() {
slog.Debug("relay ws receive timeout", "node_id", node.NodeID)
return
}
slog.Debug("relay ws receive failed", "node_id", node.NodeID, "error", err)
return
}
slog.Debug("relay ws message received", "node_id", node.NodeID, "type", message.Type)
switch message.Type {
case "ping":
if !service.SendRelayWSPong(node.NodeID) {
slog.Debug("relay ws pong enqueue failed", "node_id", node.NodeID)
}
case "pong":
slog.Debug("relay ws pong received", "node_id", node.NodeID)
default:
slog.Debug("relay ws unsupported message type", "node_id", node.NodeID, "type", message.Type)
}
}
}).ServeHTTP(c.Writer, c.Request)
}
func streamRelayWSMessages(c *gin.Context, conn *websocket.Conn, client *service.WSClient) {
for {
select {
case <-c.Request.Context().Done():
return
case <-client.Done():
return
case message, ok := <-client.Messages():
if !ok {
return
}
_ = conn.SetWriteDeadline(time.Now().Add(agentWSWriteTimeout()))
if err := websocket.JSON.Send(conn, message); err != nil {
slog.Debug("relay ws send failed", "node_id", client.ID(), "error", err)
return
}
}
}
}
@@ -0,0 +1,282 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
)
// GetTLSCertificates godoc
// @Summary List TLS certificates
// @Tags TLSCertificates
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/tls-certificates/ [get]
func GetTLSCertificates(c *gin.Context) {
certificates, err := service.ListTLSCertificates()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, certificates)
}
// GetTLSCertificate godoc
// @Summary Get TLS certificate detail
// @Tags TLSCertificates
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Certificate ID"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/{id} [get]
func GetTLSCertificate(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
certificate, err := service.GetTLSCertificate(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, certificate)
}
// GetTLSCertificateContent godoc
// @Summary Get TLS certificate PEM content
// @Tags TLSCertificates
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Certificate ID"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/{id}/content [get]
func GetTLSCertificateContent(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
content, err := service.GetTLSCertificateContent(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, content)
}
// CreateTLSCertificate godoc
// @Summary Create TLS certificate from PEM
// @Tags TLSCertificates
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Param payload body service.TLSCertificateInput true "TLS certificate payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/ [post]
func CreateTLSCertificate(c *gin.Context) {
var input service.TLSCertificateInput
if !bind.JSON(c, &input) {
return
}
certificate, err := service.CreateTLSCertificate(input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, certificate)
}
// UpdateTLSCertificate godoc
// @Summary Update TLS certificate from PEM
// @Tags TLSCertificates
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Certificate ID"
// @Param payload body service.TLSCertificateInput true "TLS certificate payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/{id}/update [post]
func UpdateTLSCertificate(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
var input service.TLSCertificateInput
if !bind.JSON(c, &input) {
return
}
certificate, err := service.UpdateTLSCertificate(id, input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, certificate)
}
// ImportTLSCertificateFile godoc
// @Summary Import TLS certificate from files
// @Tags TLSCertificates
// @Accept multipart/form-data
// @Produce json
// @Security OpenFlareTokenAuth
// @Param name formData string true "Certificate name"
// @Param remark formData string false "Remark"
// @Param cert_file formData file true "Certificate file"
// @Param key_file formData file true "Private key file"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/import-file [post]
func ImportTLSCertificateFile(c *gin.Context) {
name := c.PostForm("name")
remark := c.PostForm("remark")
certFile, err := c.FormFile("cert_file")
if err != nil {
response.RespondBadRequest(c, "缺少证书文件")
return
}
keyFile, err := c.FormFile("key_file")
if err != nil {
response.RespondBadRequest(c, "缺少私钥文件")
return
}
certificate, err := service.CreateTLSCertificateFromFiles(name, certFile, keyFile, remark)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, certificate)
}
// DeleteTLSCertificate godoc
// @Summary Delete TLS certificate
// @Tags TLSCertificates
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Certificate ID"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/{id}/delete [post]
func DeleteTLSCertificate(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
if err := service.DeleteTLSCertificate(id); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, nil)
}
// ApplyTLSCertificate godoc
// @Summary Apply TLS certificate via ACME
// @Tags TLSCertificates
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Param payload body service.TLSApplyInput true "TLS apply payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/apply [post]
func ApplyTLSCertificate(c *gin.Context) {
var input service.TLSApplyInput
if !bind.JSON(c, &input) {
return
}
certificate, err := service.ApplyTLSCertificate(input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, certificate)
}
// UpdateAcmeCertificate godoc
// @Summary Update ACME TLS certificate
// @Tags TLSCertificates
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Certificate ID"
// @Param payload body service.TLSApplyInput true "TLS apply payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/{id}/update-acme [post]
func UpdateAcmeCertificate(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
var input service.TLSApplyInput
if !bind.JSON(c, &input) {
return
}
certificate, err := service.UpdateAcmeCertificate(id, input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, certificate)
}
// ConvertTLSCertificateToAcme godoc
// @Summary Convert uploaded TLS certificate to ACME managed certificate
// @Tags TLSCertificates
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Certificate ID"
// @Param payload body service.TLSApplyInput true "TLS apply payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/{id}/convert-acme [post]
func ConvertTLSCertificateToAcme(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
var input service.TLSApplyInput
if !bind.JSON(c, &input) {
return
}
certificate, err := service.ConvertTLSCertificateToAcme(id, input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, certificate)
}
// RenewTLSCertificate godoc
// @Summary Renew TLS certificate
// @Tags TLSCertificates
// @Produce json
// @Security OpenFlareTokenAuth
// @Param id path int true "Certificate ID"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/{id}/renew [post]
func RenewTLSCertificate(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
certificate, err := service.RenewTLSCertificate(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, certificate)
}
@@ -0,0 +1,166 @@
package controller
import (
"time"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
"golang.org/x/net/websocket"
)
type confirmManualUpgradeRequest struct {
UploadToken string `json:"upload_token"`
}
type serverUpgradeRequest struct {
Channel string `json:"channel"`
}
// GetLatestRelease godoc
// @Summary Get latest GitHub release
// @Tags Update
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/update/latest-release [get]
func GetLatestRelease(c *gin.Context) {
release, err := service.GetLatestServerRelease(c.Request.Context(), c.Query("channel"))
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, release)
}
// UpgradeServer godoc
// @Summary Upgrade server binary from latest GitHub release
// @Tags Update
// @Produce json
// @Success 200 {object} map[string]interface{}
// @Router /api/update/upgrade [post]
func UpgradeServer(c *gin.Context) {
var request serverUpgradeRequest
if c.Request.ContentLength > 0 {
if err := bind.OptionalJSON(c.Request.Body, &request); err != nil {
response.RespondBadRequest(c, "无效的参数")
return
}
}
release, err := service.ScheduleServerUpgrade(request.Channel)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessWithExtras(c, release, gin.H{
"message": "服务升级任务已启动,下载完成后将自动重启。",
})
}
// StreamServerUpgradeLogs godoc
// @Summary Stream server upgrade logs over websocket
// @Tags Update
// @Router /api/update/logs/ws [get]
func StreamServerUpgradeLogs(c *gin.Context) {
websocket.Handler(func(conn *websocket.Conn) {
defer func() {
_ = conn.Close()
}()
updates, unsubscribe := service.SubscribeServerUpgradeStream()
defer unsubscribe()
heartbeatTicker := time.NewTicker(15 * time.Second)
defer heartbeatTicker.Stop()
for {
select {
case snapshot, ok := <-updates:
if !ok {
return
}
if err := websocket.JSON.Send(conn, snapshot); err != nil {
return
}
case <-heartbeatTicker.C:
if err := websocket.JSON.Send(conn, service.ServerUpgradeStreamSnapshot{}); err != nil {
return
}
case <-c.Request.Context().Done():
return
}
}
}).ServeHTTP(c.Writer, c.Request)
}
// UploadManualServerBinary godoc
// @Summary Upload server binary and inspect version before upgrade
// @Tags Update
// @Accept mpfd
// @Produce json
// @Success 200 {object} map[string]interface{}
// @Router /api/update/manual-upload [post]
func UploadManualServerBinary(c *gin.Context) {
response.RespondFailure(c, "手动升级功能已禁用")
return
//
//fileHeader, err := c.FormFile("binary")
//if err != nil {
// response.RespondFailure(c, "请先选择要上传的服务端二进制文件。")
// return
//}
//
//file, err := fileHeader.Open()
//if err != nil {
// response.RespondFailure(c, "读取上传文件失败。")
// return
//}
//defer func() {
// _ = file.Close()
//}()
//
//info, err := service.UploadManualServerBinary(c.Request.Context(), fileHeader.Filename, file)
//if err != nil {
// response.RespondFailure(c, err.Error())
// return
//}
//
//message := strings.TrimSpace(info.ComparisonMessage)
//if message == "" {
// message = "已完成上传并检查升级包版本。"
//}
//
//response.RespondSuccessWithExtras(c, info, gin.H{
// "message": message,
//})
}
// ConfirmManualServerUpgrade godoc
// @Summary Confirm upgrade with previously uploaded server binary
// @Tags Update
// @Accept json
// @Produce json
// @Success 200 {object} map[string]interface{}
// @Router /api/update/manual-upgrade [post]
func ConfirmManualServerUpgrade(c *gin.Context) {
response.RespondFailure(c, "手动升级功能已禁用")
return
//
//var request confirmManualUpgradeRequest
//if !bind.JSON(c, &request) {
// return
//}
//
//info, err := service.ConfirmManualServerUpgrade(request.UploadToken)
//if err != nil {
// response.RespondFailure(c, err.Error())
// return
//}
//
//response.RespondSuccessWithExtras(c, info, gin.H{
// "message": "服务升级任务已启动,确认无误后将自动重启。",
//})
}
@@ -0,0 +1,25 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
)
// SyncUptimeKuma godoc
// @Summary Manually trigger Uptime Kuma sync
// @Tags UptimeKuma
// @Accept json
// @Produce json
// @Security OpenFlareTokenAuth
// @Success 200 {object} map[string]interface{}
// @Router /api/uptimekuma/sync [post]
func SyncUptimeKuma(c *gin.Context) {
err := service.SyncToUptimeKuma()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "同步成功")
}
@@ -0,0 +1,415 @@
package controller
import (
"strconv"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/middleware"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/rain-kl/openflare/openflare-server/internal/utils/security"
"github.com/rain-kl/openflare/openflare-server/internal/utils/validation"
"github.com/gin-gonic/gin"
)
type LoginRequest struct {
Username string `json:"username"`
Password string `json:"password"`
}
func Login(c *gin.Context) {
if !common.PasswordLoginEnabled {
response.RespondFailure(c, "管理员关闭了密码登录")
return
}
var loginRequest LoginRequest
if !bind.JSON(c, &loginRequest) {
return
}
username := loginRequest.Username
password := loginRequest.Password
if username == "" || password == "" {
response.RespondFailure(c, "无效的参数")
return
}
user := model.User{
Username: username,
Password: password,
}
err := user.ValidateAndFill()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
setupLogin(&user, c)
}
// setup token and then return user info
func setLoginToken(user *model.User) (*model.User, error) {
// Generate a signed JWT using gin-jwt middleware
tokenString, _, err := middleware.JWTMiddleware.TokenGenerator(user)
if err != nil {
return nil, err
}
// Persist JWT in DB so we can invalidate it on logout
if err := model.DB.Model(user).Update("token", tokenString).Error; err != nil {
return nil, err
}
cleanUser := &model.User{
Id: user.Id,
Username: user.Username,
DisplayName: user.DisplayName,
Role: user.Role,
Status: user.Status,
Token: tokenString,
}
return cleanUser, nil
}
func setupLogin(user *model.User, c *gin.Context) {
cleanUser, err := setLoginToken(user)
if err != nil {
response.RespondFailure(c, "无法保存会话信息,请重试")
return
}
response.RespondSuccess(c, *cleanUser)
}
func Logout(c *gin.Context) {
token := c.GetHeader("OpenFlare-Token")
if token != "" {
user := model.ValidateUserToken(token)
if user != nil && user.Id != 0 {
if err := model.DB.Model(user).Update("token", "").Error; err != nil {
response.RespondFailure(c, err.Error())
return
}
}
}
response.RespondSuccessMessage(c, "")
}
func currentUserFromOpenFlareToken(c *gin.Context) *model.User {
token := c.GetHeader("OpenFlare-Token")
if token == "" {
return nil
}
return model.ValidateUserToken(token)
}
func Register(c *gin.Context) {
response.RespondFailure(c, "非法请求")
}
func GetAllUsers(c *gin.Context) {
p, _ := strconv.Atoi(c.Query("p"))
if p < 0 {
p = 0
}
users, err := model.GetAllUsers(p*common.ItemsPerPage, common.ItemsPerPage)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, users)
}
func SearchUsers(c *gin.Context) {
keyword := c.Query("keyword")
users, err := model.SearchUsers(keyword)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, users)
}
func GetUser(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
user, err := model.GetUserById(int(id), false)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
myRole := c.GetInt("role")
if myRole <= user.Role {
response.RespondFailure(c, "无权获取同级或更高等级用户的信息")
return
}
response.RespondSuccess(c, user)
}
func GenerateToken(c *gin.Context) {
id := c.GetInt("id")
user, err := model.GetUserById(id, true)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
// Generate a fresh JWT for the user
tokenString, _, err := middleware.JWTMiddleware.TokenGenerator(user)
if err != nil {
response.RespondFailure(c, "生成 Token 失败: "+err.Error())
return
}
user.Token = tokenString
if err := user.Update(false); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, user.Token)
}
func GetSelf(c *gin.Context) {
id := c.GetInt("id")
user, err := model.GetUserById(id, false)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, user)
}
func UpdateUser(c *gin.Context) {
var updatedUser model.User
if !bind.JSON(c, &updatedUser) {
return
}
if updatedUser.Id == 0 {
response.RespondFailure(c, "无效的参数")
return
}
if updatedUser.Password == "" {
updatedUser.Password = "$I_LOVE_U" // make Validator happy :)
}
if err := validation.Validate.Struct(&updatedUser); err != nil {
response.RespondFailure(c, "输入不合法 "+err.Error())
return
}
originUser, err := model.GetUserById(updatedUser.Id, false)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
myRole := c.GetInt("role")
if myRole <= originUser.Role {
response.RespondFailure(c, "无权更新同权限等级或更高权限等级的用户信息")
return
}
if myRole <= updatedUser.Role {
response.RespondFailure(c, "无权将其他用户权限等级提升到大于等于自己的权限等级")
return
}
if updatedUser.Password == "$I_LOVE_U" {
updatedUser.Password = "" // rollback to what it should be
}
updatePassword := updatedUser.Password != ""
if err := updatedUser.Update(updatePassword); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
func UpdateSelf(c *gin.Context) {
var user model.User
if !bind.JSON(c, &user) {
return
}
if user.Password == "" {
user.Password = "$I_LOVE_U" // make Validator happy :)
}
if err := validation.Validate.Struct(&user); err != nil {
response.RespondFailure(c, "输入不合法 "+err.Error())
return
}
cleanUser := model.User{
Id: c.GetInt("id"),
Username: user.Username,
Password: user.Password,
DisplayName: user.DisplayName,
}
if user.Password == "$I_LOVE_U" {
user.Password = "" // rollback to what it should be
cleanUser.Password = ""
}
updatePassword := user.Password != ""
if err := cleanUser.Update(updatePassword); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
func DeleteUser(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
originUser, err := model.GetUserById(int(id), false)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
myRole := c.GetInt("role")
if myRole <= originUser.Role {
response.RespondFailure(c, "无权删除同权限等级或更高权限等级的用户")
return
}
err = model.DeleteUserById(int(id))
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
func DeleteSelf(c *gin.Context) {
id := c.GetInt("id")
err := model.DeleteUserById(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
func CreateUser(c *gin.Context) {
var user model.User
if !bind.JSON(c, &user) {
return
}
if user.Username == "" || user.Password == "" {
response.RespondFailure(c, "无效的参数")
return
}
if user.DisplayName == "" {
user.DisplayName = user.Username
}
myRole := c.GetInt("role")
if user.Role >= myRole {
response.RespondFailure(c, "无法创建权限大于等于自己的用户")
return
}
// Even for admin users, we cannot fully trust them!
cleanUser := model.User{
Username: user.Username,
Password: user.Password,
DisplayName: user.DisplayName,
}
if err := cleanUser.Insert(); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
type ManageRequest struct {
Username string `json:"username"`
Action string `json:"action"`
}
// ManageUser Only admin user can do this
func ManageUser(c *gin.Context) {
var req ManageRequest
if !bind.JSON(c, &req) {
return
}
user := model.User{
Username: req.Username,
}
// Fill attributes
model.DB.Where(&user).First(&user)
if user.Id == 0 {
response.RespondFailure(c, "用户不存在")
return
}
myRole := c.GetInt("role")
if myRole <= user.Role && myRole != common.RoleRootUser {
response.RespondFailure(c, "无权更新同权限等级或更高权限等级的用户信息")
return
}
switch req.Action {
case "disable":
user.Status = common.UserStatusDisabled
if user.Role == common.RoleRootUser {
response.RespondFailure(c, "无法禁用超级管理员用户")
return
}
case "enable":
user.Status = common.UserStatusEnabled
case "delete":
if user.Role == common.RoleRootUser {
response.RespondFailure(c, "无法删除超级管理员用户")
return
}
if err := user.Delete(); err != nil {
response.RespondFailure(c, err.Error())
return
}
case "promote":
if myRole != common.RoleRootUser {
response.RespondFailure(c, "普通管理员用户无法提升其他用户为管理员")
return
}
if user.Role >= common.RoleAdminUser {
response.RespondFailure(c, "该用户已经是管理员")
return
}
user.Role = common.RoleAdminUser
case "demote":
if user.Role == common.RoleRootUser {
response.RespondFailure(c, "无法降级超级管理员用户")
return
}
if user.Role == common.RoleCommonUser {
response.RespondFailure(c, "该用户已经是普通用户")
return
}
user.Role = common.RoleCommonUser
}
if err := user.Update(false); err != nil {
response.RespondFailure(c, err.Error())
return
}
clearUser := model.User{
Role: user.Role,
Status: user.Status,
}
response.RespondSuccess(c, clearUser)
}
func EmailBind(c *gin.Context) {
email := c.Query("email")
code := c.Query("code")
if !security.VerifyCodeWithKey(email, code, security.EmailVerificationPurpose) {
response.RespondFailure(c, "验证码错误或已过期")
return
}
id := c.GetInt("id")
user := model.User{
Id: id,
}
err := user.FillUserById()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
user.Email = email
// no need to check if this email already taken, because we have used verification code to check it
err = user.Update(false)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
+225
View File
@@ -0,0 +1,225 @@
package controller
import (
"strconv"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/controller/bind"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
)
type wafIDsRequest struct {
IDs []uint `json:"ids"`
}
func ListWAFRuleGroups(c *gin.Context) {
groups, err := service.ListWAFRuleGroups()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, groups)
}
func GetWAFRuleGroup(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
group, err := service.GetWAFRuleGroup(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, group)
}
func CreateWAFRuleGroup(c *gin.Context) {
var input service.WAFRuleGroupInput
if !bind.JSON(c, &input) {
return
}
group, err := service.CreateWAFRuleGroup(input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, group)
}
func UpdateWAFRuleGroup(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
var input service.WAFRuleGroupInput
if !bind.JSON(c, &input) {
return
}
group, err := service.UpdateWAFRuleGroup(id, input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, group)
}
func DeleteWAFRuleGroup(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
if err := service.DeleteWAFRuleGroup(id); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
func ReplaceWAFRuleGroupSites(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
var request wafIDsRequest
if !bind.JSON(c, &request) {
return
}
group, err := service.ReplaceWAFRuleGroupSites(id, request.IDs)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, group)
}
func GetWAFSiteRuleGroups(c *gin.Context) {
routeID, ok := parseUintPathParam(c, "route_id")
if !ok {
return
}
view, err := service.GetWAFSiteRuleGroups(routeID)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, view)
}
func ReplaceWAFSiteRuleGroups(c *gin.Context) {
routeID, ok := parseUintPathParam(c, "route_id")
if !ok {
return
}
var request wafIDsRequest
if !bind.JSON(c, &request) {
return
}
view, err := service.ReplaceWAFSiteRuleGroups(routeID, request.IDs)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, view)
}
func ListWAFIPGroups(c *gin.Context) {
groups, err := service.ListWAFIPGroups()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, groups)
}
func GetWAFIPGroup(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
group, err := service.GetWAFIPGroup(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, group)
}
func CreateWAFIPGroup(c *gin.Context) {
var input service.WAFIPGroupInput
if !bind.JSON(c, &input) {
return
}
group, err := service.CreateWAFIPGroup(input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, group)
}
func UpdateWAFIPGroup(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
var input service.WAFIPGroupInput
if !bind.JSON(c, &input) {
return
}
group, err := service.UpdateWAFIPGroup(id, input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, group)
}
func DeleteWAFIPGroup(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
if err := service.DeleteWAFIPGroup(id); err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
}
func SyncWAFIPGroup(c *gin.Context) {
id, ok := bind.IDParam(c)
if !ok {
return
}
result, err := service.SyncWAFIPGroup(id)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, result)
}
func TestWAFIPGroupAutoConfig(c *gin.Context) {
var input service.WAFIPGroupAutoTestInput
if !bind.JSON(c, &input) {
return
}
result, err := service.TestWAFIPGroupAutoConfig(input)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, result)
}
func parseUintPathParam(c *gin.Context, name string) (uint, bool) {
id, err := strconv.ParseUint(c.Param(name), 10, 64)
if err != nil || id == 0 {
response.RespondBadRequest(c, "invalid id")
return 0, false
}
return uint(id), true
}
@@ -0,0 +1,125 @@
package controller
import (
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"time"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/gin-gonic/gin"
)
type wechatLoginResponse struct {
Success bool `json:"success"`
Message string `json:"message"`
Data string `json:"data"`
}
func getWeChatIdByCode(code string) (string, error) {
if code == "" {
return "", errors.New("无效的参数")
}
req, err := http.NewRequest("GET", fmt.Sprintf("%s/api/wechat/user?code=%s", common.WeChatServerAddress, code), nil)
if err != nil {
return "", err
}
req.Header.Set("Authorization", common.WeChatServerToken)
client := http.Client{
Timeout: 5 * time.Second,
}
httpResponse, err := client.Do(req)
if err != nil {
return "", err
}
defer func(Body io.ReadCloser) {
err := Body.Close()
if err != nil {
slog.Error("Failed to close response body", "error", err)
}
}(httpResponse.Body)
var res wechatLoginResponse
err = json.NewDecoder(httpResponse.Body).Decode(&res)
if err != nil {
return "", err
}
if !res.Success {
return "", errors.New(res.Message)
}
if res.Data == "" {
return "", errors.New("验证码错误或已过期")
}
return res.Data, nil
}
func WeChatAuth(c *gin.Context) {
if !common.WeChatAuthEnabled {
response.RespondFailure(c, "管理员未开启通过微信登录以及注册")
return
}
code := c.Query("code")
wechatId, err := getWeChatIdByCode(code)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
user := model.User{
WeChatId: wechatId,
}
if model.IsWeChatIdAlreadyTaken(wechatId) {
err := user.FillUserByWeChatId()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
} else {
response.RespondFailure(c, "管理员关闭了新用户注册")
return
}
if user.Status != common.UserStatusEnabled {
response.RespondFailure(c, "用户已被封禁")
return
}
setupLogin(&user, c)
}
func WeChatBind(c *gin.Context) {
if !common.WeChatAuthEnabled {
response.RespondFailure(c, "管理员未开启通过微信登录以及注册")
return
}
code := c.Query("code")
wechatId, err := getWeChatIdByCode(code)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
if model.IsWeChatIdAlreadyTaken(wechatId) {
response.RespondFailure(c, "该微信账号已被绑定")
return
}
id := c.GetInt("id")
user := model.User{
Id: id,
}
err = user.FillUserById()
if err != nil {
response.RespondFailure(c, err.Error())
return
}
user.WeChatId = wechatId
err = user.Update(false)
if err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccessMessage(c, "")
return
}
+44
View File
@@ -0,0 +1,44 @@
package job
import (
"log/slog"
"github.com/robfig/cron/v3"
)
var cronRunner *cron.Cron
func InitCronJobs() {
cronRunner = cron.New()
// Register SSL renew job
_, err := cronRunner.AddJob("0 0 * * *", &SSLRenewJob{})
if err != nil {
slog.Error("failed to register SSL renew cron job", "error", err)
} else {
slog.Info("registered SSL renew cron job")
}
_, err = cronRunner.AddJob("@every 5m", &WAFIPGroupSyncJob{})
if err != nil {
slog.Error("failed to register WAF IP group sync cron job", "error", err)
} else {
slog.Info("registered WAF IP group sync cron job")
}
// Register Uptime Kuma sync job (check every minute)
_, err = cronRunner.AddJob("* * * * *", &UptimeKumaSyncJob{})
if err != nil {
slog.Error("failed to register Uptime Kuma sync cron job", "error", err)
} else {
slog.Info("registered Uptime Kuma sync cron job")
}
cronRunner.Start()
}
func StopCronJobs() {
if cronRunner != nil {
cronRunner.Stop()
}
}
@@ -0,0 +1,44 @@
package job
import (
"log/slog"
"time"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/rain-kl/openflare/openflare-server/internal/service"
)
type SSLRenewJob struct {
}
func (j *SSLRenewJob) Run() {
slog.Info("The scheduled certificate update task is currently in progress ...")
certificates, err := model.ListTLSCertificates()
if err != nil {
slog.Error("failed to list certificates in SSL renew job", "error", err)
return
}
now := time.Now()
for _, cert := range certificates {
if !cert.AutoRenew || cert.Provider != "acme" || cert.ApplyStatus == "applying" {
continue
}
sub := cert.NotAfter.Sub(now)
// Expiring in less than 7 days (7 * 24 hours)
if sub.Hours() < 168 {
slog.Info("Update the SSL certificate for the domain", "domain", cert.PrimaryDomain)
// Invoke renew process (async go-routine handles Lego inside)
_, err := service.RenewTLSCertificate(cert.ID)
if err != nil {
slog.Error("Failed to update the SSL certificate", "domain", cert.PrimaryDomain, "error", err)
continue
}
slog.Info("Triggered the SSL certificate renew for domain", "domain", cert.PrimaryDomain)
}
}
slog.Info("The scheduled certificate update task has completed")
}
@@ -0,0 +1,44 @@
package job
import (
"log/slog"
"sync"
"time"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/service"
)
var lastUptimeKumaSyncTime time.Time
var uptimeKumaSyncMutex sync.Mutex
type UptimeKumaSyncJob struct{}
func (j *UptimeKumaSyncJob) Run() {
if !common.UptimeKumaEnabled {
return
}
interval := common.UptimeKumaSyncInterval
if interval <= 0 {
interval = 5
}
if time.Since(lastUptimeKumaSyncTime) < time.Duration(interval)*time.Minute {
return
}
if !uptimeKumaSyncMutex.TryLock() {
slog.Warn("Uptime Kuma sync job is already running, skipping this scheduled run")
return
}
defer uptimeKumaSyncMutex.Unlock()
slog.Info("Starting scheduled Uptime Kuma sync")
if err := service.SyncToUptimeKuma(); err != nil {
slog.Error("Uptime Kuma sync failed", "error", err)
} else {
lastUptimeKumaSyncTime = time.Now()
slog.Info("Uptime Kuma sync completed successfully")
}
}
@@ -0,0 +1,15 @@
package job
import (
"log/slog"
"github.com/rain-kl/openflare/openflare-server/internal/service"
)
type WAFIPGroupSyncJob struct{}
func (j *WAFIPGroupSyncJob) Run() {
if err := service.SyncDueWAFIPGroups(); err != nil {
slog.Error("failed to sync due waf ip groups", "error", err)
}
}
@@ -0,0 +1,39 @@
package middleware
import (
"github.com/gin-gonic/gin"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/service"
)
func AgentAuth() func(c *gin.Context) {
return func(c *gin.Context) {
token := c.GetHeader("X-Agent-Token")
node, err := service.AuthenticateAccessToken(token)
if err != nil {
response.RespondUnauthorized(c, "无权进行此操作,Agent Token 无效")
c.Abort()
return
}
c.Set("agent_node", node)
c.Next()
}
}
func AgentRegisterAuth() func(c *gin.Context) {
return func(c *gin.Context) {
token := c.GetHeader("X-Agent-Token")
if node, err := service.AuthenticateAccessToken(token); err == nil {
c.Set("agent_node", node)
c.Next()
return
}
if err := service.ValidateDiscoveryToken(token); err != nil {
response.RespondUnauthorized(c, "无权进行此操作,注册 Token 无效")
c.Abort()
return
}
c.Set("discovery_enabled", true)
c.Next()
}
}
@@ -0,0 +1,101 @@
package middleware
import (
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/model"
jwt "github.com/appleboy/gin-jwt/v2"
"github.com/gin-gonic/gin"
)
const OpenFlareTokenHeader = "OpenFlare-Token"
func authHelper(c *gin.Context, minRole int) {
tokenStr := c.GetHeader(OpenFlareTokenHeader)
if tokenStr == "" {
response.RespondUnauthorized(c, "无权进行此操作,未登录或 token 无效")
c.Abort()
return
}
token, err := JWTMiddleware.ParseTokenString(tokenStr)
if err != nil {
response.RespondUnauthorized(c, "无权进行此操作,token 无效: "+err.Error())
c.Abort()
return
}
claims := jwt.ExtractClaimsFromToken(token)
id, ok := claims["id"].(float64)
if !ok {
response.RespondUnauthorized(c, "无权进行此操作,token 格式错误")
c.Abort()
return
}
dbUser := &model.User{}
dbErr := model.DB.Select([]string{"id", "username", "display_name", "role", "status", "token"}).
First(dbUser, "id = ?", int(id)).Error
if dbErr != nil || dbUser.Username == "" {
response.RespondUnauthorized(c, "无权进行此操作,用户不存在")
c.Abort()
return
}
if dbUser.Token != tokenStr {
response.RespondUnauthorized(c, "无权进行此操作,token 已失效或已登出")
c.Abort()
return
}
if dbUser.Status == common.UserStatusDisabled {
response.RespondFailure(c, "用户已被封禁")
c.Abort()
return
}
if int(dbUser.Role) < minRole {
response.RespondFailure(c, "无权进行此操作,权限不足")
c.Abort()
return
}
c.Set("username", dbUser.Username)
c.Set("role", dbUser.Role)
c.Set("id", dbUser.Id)
c.Set("authByToken", true)
c.Next()
}
func UserAuth() func(c *gin.Context) {
return func(c *gin.Context) {
authHelper(c, common.RoleCommonUser)
}
}
func AdminAuth() func(c *gin.Context) {
return func(c *gin.Context) {
authHelper(c, common.RoleAdminUser)
}
}
func RootAuth() func(c *gin.Context) {
return func(c *gin.Context) {
authHelper(c, common.RoleRootUser)
}
}
// NoTokenAuth is kept as a compatibility no-op because admin APIs now always use OPENFLARE_TOKEN.
func NoTokenAuth() func(c *gin.Context) {
return func(c *gin.Context) {
c.Next()
}
}
// TokenOnlyAuth is kept as a compatibility no-op because admin APIs now always use OPENFLARE_TOKEN.
func TokenOnlyAuth() func(c *gin.Context) {
return func(c *gin.Context) {
c.Next()
}
}
@@ -0,0 +1,37 @@
package middleware
import (
"path"
"strings"
"github.com/gin-gonic/gin"
)
func Cache() func(c *gin.Context) {
return func(c *gin.Context) {
requestPath := c.Request.URL.Path
switch {
case strings.HasPrefix(requestPath, "/_next/static/"):
c.Header("Cache-Control", "public, max-age=31536000, immutable")
case isStaticPublicAsset(requestPath):
c.Header("Cache-Control", "public, max-age=86400")
default:
c.Header("Cache-Control", "no-store, no-cache, must-revalidate")
c.Header("Pragma", "no-cache")
c.Header("Expires", "0")
}
c.Next()
}
}
func isStaticPublicAsset(requestPath string) bool {
ext := strings.ToLower(path.Ext(requestPath))
switch ext {
case ".ico", ".png", ".jpg", ".jpeg", ".gif", ".svg", ".webp", ".css", ".js":
return true
default:
return false
}
}
@@ -0,0 +1,14 @@
package middleware
import (
"github.com/gin-gonic/gin"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/service"
)
// CapAuth wraps the core Cap middleware with OpenFlare's dynamic CapLoginEnabled configuration switch
func CapAuth(scope string) gin.HandlerFunc {
return service.CapManager.VerifyMiddleware(scope, func() bool {
return common.CapLoginEnabled
})
}
@@ -0,0 +1,27 @@
package middleware
import (
"strings"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/gin-contrib/cors"
"github.com/gin-gonic/gin"
)
func CORS() gin.HandlerFunc {
config := cors.DefaultConfig()
config.AllowCredentials = true
config.AllowHeaders = []string{"Origin", "Content-Length", "Content-Type", "Authorization", "OpenFlare-Token", "X-Agent-Token", "Accept"}
config.AllowOriginFunc = func(origin string) bool {
serverAddr := strings.TrimRight(common.ServerAddress, "/")
if serverAddr == "" {
return true
}
if origin == serverAddr {
return true
}
return false
}
return cors.New(config)
}
@@ -0,0 +1,72 @@
package middleware
import (
"log"
"time"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/model"
jwt "github.com/appleboy/gin-jwt/v2"
"github.com/gin-gonic/gin"
)
var JWTMiddleware *jwt.GinJWTMiddleware
// jwtSigningKey returns JWT_SECRET when set, falling back to SESSION_SECRET
// for backward compatibility with deployments that only configure SESSION_SECRET.
func jwtSigningKey() []byte {
if common.JWTSecret != "" {
return []byte(common.JWTSecret)
}
return []byte(common.SessionSecret)
}
func InitJWTMiddleware() {
var err error
JWTMiddleware, err = jwt.New(&jwt.GinJWTMiddleware{
Realm: "openflare",
Key: jwtSigningKey(),
Timeout: 24 * time.Hour,
MaxRefresh: 24 * time.Hour,
IdentityKey: "identity",
PayloadFunc: func(data interface{}) jwt.MapClaims {
if v, ok := data.(*model.User); ok {
return jwt.MapClaims{
"id": v.Id,
"username": v.Username,
"role": v.Role,
}
}
return jwt.MapClaims{}
},
IdentityHandler: func(c *gin.Context) interface{} {
claims := jwt.ExtractClaims(c)
id, ok := claims["id"].(float64)
if !ok {
return nil
}
username, _ := claims["username"].(string)
role, _ := claims["role"].(float64)
return &model.User{
Id: int(id),
Username: username,
Role: int(role),
}
},
Authorizator: func(data interface{}, c *gin.Context) bool {
return data != nil
},
Unauthorized: func(c *gin.Context, code int, message string) {
response.RespondErrorWithStatus(c, code, "无权进行此操作,未登录或 token 无效: "+message)
},
TokenLookup: "header: OpenFlare-Token",
TokenHeadName: "", // Empty for raw token value directly
SendCookie: false,
})
if err != nil {
log.Fatalf("JWT Init Error: %s", err.Error())
}
}
@@ -0,0 +1,98 @@
package middleware
import (
"context"
"log/slog"
"net/http"
"time"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/utils/ratelimit"
"github.com/gin-gonic/gin"
)
var timeFormat = "2006-01-02T15:04:05.000Z"
var inMemoryRateLimiter ratelimit.InMemoryRateLimiter
func redisRateLimiter(c *gin.Context, maxRequestNum int, duration int64, mark string) {
ctx := context.Background()
rdb := common.RDB
key := "rateLimit:" + mark + c.ClientIP()
listLength, err := rdb.LLen(ctx, key).Result()
if err != nil {
slog.Error("redis rate limiter llen failed", "error", err)
c.Status(http.StatusInternalServerError)
c.Abort()
return
}
if listLength < int64(maxRequestNum) {
rdb.LPush(ctx, key, time.Now().Format(timeFormat))
rdb.Expire(ctx, key, common.RateLimitKeyExpirationDuration)
} else {
oldTimeStr, _ := rdb.LIndex(ctx, key, -1).Result()
oldTime, err := time.Parse(timeFormat, oldTimeStr)
if err != nil {
slog.Error("parse redis rate limiter old timestamp failed", "error", err)
c.Status(http.StatusInternalServerError)
c.Abort()
return
}
nowTimeStr := time.Now().Format(timeFormat)
nowTime, err := time.Parse(timeFormat, nowTimeStr)
if err != nil {
slog.Error("parse redis rate limiter current timestamp failed", "error", err)
c.Status(http.StatusInternalServerError)
c.Abort()
return
}
// time.Since will return negative number!
// See: https://stackoverflow.com/questions/50970900/why-is-time-since-returning-negative-durations-on-windows
if int64(nowTime.Sub(oldTime).Seconds()) < duration {
rdb.Expire(ctx, key, common.RateLimitKeyExpirationDuration)
c.Status(http.StatusTooManyRequests)
c.Abort()
return
}
rdb.LPush(ctx, key, time.Now().Format(timeFormat))
rdb.LTrim(ctx, key, 0, int64(maxRequestNum-1))
rdb.Expire(ctx, key, common.RateLimitKeyExpirationDuration)
}
}
func memoryRateLimiter(c *gin.Context, maxRequestNum int, duration int64, mark string) {
key := mark + c.ClientIP()
if !inMemoryRateLimiter.Request(key, maxRequestNum, duration) {
c.Status(http.StatusTooManyRequests)
c.Abort()
return
}
}
func rateLimitFactory(maxRequestNum int, duration int64, mark string) func(c *gin.Context) {
if common.RedisEnabled {
return func(c *gin.Context) {
redisRateLimiter(c, maxRequestNum, duration, mark)
}
}
// It's safe to call multi times.
inMemoryRateLimiter.Init(common.RateLimitKeyExpirationDuration)
return func(c *gin.Context) {
memoryRateLimiter(c, maxRequestNum, duration, mark)
}
}
func GlobalWebRateLimit() func(c *gin.Context) {
return rateLimitFactory(common.GlobalWebRateLimitNum, common.GlobalWebRateLimitDuration, "GW")
}
func GlobalAPIRateLimit() func(c *gin.Context) {
return rateLimitFactory(common.GlobalApiRateLimitNum, common.GlobalApiRateLimitDuration, "GA")
}
func CriticalRateLimit() func(c *gin.Context) {
return rateLimitFactory(common.CriticalRateLimitNum, common.CriticalRateLimitDuration, "CT")
}
@@ -0,0 +1,28 @@
package middleware
import (
"github.com/gin-gonic/gin"
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/service"
)
// RelayAuth authenticates Relay requests using the shared agent token,
// and verifies the node is a tunnel_relay type.
func RelayAuth() func(c *gin.Context) {
return func(c *gin.Context) {
token := c.GetHeader("X-Agent-Token")
node, err := service.AuthenticateAccessToken(token)
if err != nil {
response.RespondUnauthorized(c, "无权进行此操作,Agent Token 无效")
c.Abort()
return
}
if node.NodeType != "tunnel_relay" {
response.RespondForbidden(c, "此节点不是 TunnelRelay 类型")
c.Abort()
return
}
c.Set("relay_node", node)
c.Next()
}
}
@@ -0,0 +1,30 @@
package middleware
import (
"github.com/rain-kl/openflare/openflare-server/internal/common/response"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-gonic/gin"
)
// TunnelAuth authenticates OpenFlared client requests using the per-node
// tunnel_token carried in the X-Tunnel-Token header, and verifies the node is
// of the tunnel_client type.
func TunnelAuth() func(c *gin.Context) {
return func(c *gin.Context) {
token := c.GetHeader("X-Tunnel-Token")
node, err := service.AuthenticateAccessToken(token)
if err != nil {
response.RespondUnauthorized(c, "无权进行此操作,Tunnel Token 无效")
c.Abort()
return
}
if node.NodeType != "tunnel_client" {
response.RespondForbidden(c, "此节点不是 TunnelClient 类型")
c.Abort()
return
}
c.Set("flared_node", node)
c.Next()
}
}
@@ -0,0 +1,41 @@
package model
import "time"
type AcmeAccount struct {
ID uint `json:"id" gorm:"primaryKey"`
Email string `json:"email" gorm:"size:255"`
URL string `json:"url" gorm:"size:255"`
PrivateKey string `json:"-" gorm:"type:text;not null"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func GetAcmeAccountByID(id uint) (*AcmeAccount, error) {
account := &AcmeAccount{}
err := DB.First(account, id).Error
return account, err
}
func GetDefaultAcmeAccount() (*AcmeAccount, error) {
account := &AcmeAccount{}
err := DB.Order("id asc").First(account).Error
if err != nil {
// Auto-create a default account placeholder if none exists
account.Email = "admin@openflare.dev"
err = DB.Create(account).Error
}
return account, err
}
func (account *AcmeAccount) Insert() error {
return DB.Create(account).Error
}
func (account *AcmeAccount) Update() error {
return DB.Save(account).Error
}
func (account *AcmeAccount) Delete() error {
return DB.Delete(account).Error
}
@@ -0,0 +1,81 @@
package model
import (
"time"
"gorm.io/gorm"
)
type ApplyLogQuery struct {
NodeID string
PageNo int
PageSize int
}
type ApplyLog struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
Version string `json:"version" gorm:"size:32;not null"`
Result string `json:"result" gorm:"size:32;not null"`
Message string `json:"message" gorm:"type:text"`
Checksum string `json:"checksum" gorm:"size:64;not null;default:''"`
MainConfigChecksum string `json:"main_config_checksum" gorm:"size:64;not null;default:''"`
RouteConfigChecksum string `json:"route_config_checksum" gorm:"size:64;not null;default:''"`
SupportFileCount int `json:"support_file_count" gorm:"not null;default:0"`
CreatedAt time.Time `json:"created_at"`
}
func ListApplyLogs(query ApplyLogQuery) (logs []*ApplyLog, err error) {
db := DB.Order("id desc")
if query.NodeID != "" {
db = db.Where("node_id = ?", query.NodeID)
}
if query.PageSize > 0 {
offset := 0
if query.PageNo > 1 {
offset = (query.PageNo - 1) * query.PageSize
}
db = db.Limit(query.PageSize).Offset(offset)
}
err = db.Find(&logs).Error
return logs, err
}
func CountApplyLogs(nodeID string) (total int64, err error) {
query := DB.Model(&ApplyLog{})
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
err = query.Count(&total).Error
return total, err
}
func GetLatestApplyLogsByNodeIDs(nodeIDs []string) (map[string]*ApplyLog, error) {
result := make(map[string]*ApplyLog)
if len(nodeIDs) == 0 {
return result, nil
}
var logs []*ApplyLog
subQuery := DB.Model(&ApplyLog{}).
Select("MAX(id) AS id").
Where("node_id IN ?", nodeIDs).
Group("node_id")
if err := DB.Where("id IN (?)", subQuery).Find(&logs).Error; err != nil {
return nil, err
}
for _, log := range logs {
result[log.NodeID] = log
}
return result, nil
}
func DeleteAllApplyLogs() (deleted int64, err error) {
result := DB.Session(&gorm.Session{AllowGlobalUpdate: true}).Delete(&ApplyLog{})
return result.RowsAffected, result.Error
}
func DeleteApplyLogsBefore(before time.Time) (deleted int64, err error) {
result := DB.Where("created_at < ?", before).Delete(&ApplyLog{})
return result.RowsAffected, result.Error
}
@@ -0,0 +1,279 @@
package model
import (
"errors"
"regexp"
"strings"
"time"
"github.com/rain-kl/openflare/pkg/utils"
"gorm.io/gorm"
)
const (
AuthSourceTypeGitHub = "github"
AuthSourceTypeOIDC = "oidc"
)
var authSourceNamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_-]{0,79}$`)
type AuthSource struct {
ID uint `json:"id"`
Name string `json:"name" gorm:"uniqueIndex;size:80;not null"`
Type string `json:"type" gorm:"index;size:20;not null"`
DisplayName string `json:"display_name" gorm:"size:100"`
IsActive bool `json:"is_active" gorm:"index;not null;default:false"`
ClientID string `json:"client_id" gorm:"column:client_id;size:255"`
ClientSecret string `json:"-" gorm:"column:client_secret;size:1024"`
OpenIDDiscoveryURL string `json:"openid_discovery_url" gorm:"column:openid_discovery_url;size:1024"`
Scopes string `json:"scopes" gorm:"size:255"`
IconURL string `json:"icon_url" gorm:"size:1024"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
ClientSecretConfigured bool `json:"client_secret_configured" gorm:"-"`
}
type ExternalAccount struct {
ID uint `json:"id"`
AuthSourceID uint `json:"auth_source_id" gorm:"uniqueIndex:idx_external_account_source_external;index;not null"`
UserID int `json:"user_id" gorm:"index;not null"`
ExternalID string `json:"external_id" gorm:"uniqueIndex:idx_external_account_source_external;size:255;not null"`
ExternalUsername string `json:"external_username" gorm:"size:255"`
Email string `json:"email" gorm:"size:255"`
AuthSource AuthSource `json:"-" gorm:"constraint:OnDelete:CASCADE"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type ExternalAccountView struct {
ID uint `json:"id"`
AuthSourceID uint `json:"auth_source_id"`
AuthSourceName string `json:"auth_source_name"`
AuthSourceType string `json:"auth_source_type"`
AuthSourceLabel string `json:"auth_source_label"`
ExternalUsername string `json:"external_username"`
Email string `json:"email"`
CreatedAt time.Time `json:"created_at"`
}
func (source *AuthSource) Normalize() {
source.Type = strings.ToLower(source.Type)
utils.TrimStringFields(
&source.Name,
&source.Type,
&source.DisplayName,
&source.ClientID,
&source.ClientSecret,
&source.OpenIDDiscoveryURL,
&source.Scopes,
&source.IconURL,
)
if source.DisplayName == "" {
source.DisplayName = source.Name
}
if source.Type == AuthSourceTypeOIDC && source.Scopes == "" {
source.Scopes = "openid profile email"
}
if source.Type == AuthSourceTypeGitHub && source.Scopes == "" {
source.Scopes = "user:email"
}
}
func (source *AuthSource) Validate() error {
source.Normalize()
if source.Name == "" {
return errors.New("认证源名称不能为空")
}
if !authSourceNamePattern.MatchString(source.Name) {
return errors.New("认证源名称只能包含字母、数字、短横线或下划线,且必须以字母或数字开头")
}
switch source.Type {
case AuthSourceTypeGitHub:
case AuthSourceTypeOIDC:
if source.OpenIDDiscoveryURL == "" {
return errors.New("OIDC 认证源必须配置 Discovery URL")
}
default:
return errors.New("认证源类型仅支持 github 或 oidc")
}
if source.IsActive {
if source.ClientID == "" || source.ClientSecret == "" {
return errors.New("启用认证源前必须配置 Client ID 和 Client Secret")
}
}
return nil
}
func (source *AuthSource) Sanitize() {
source.ClientSecretConfigured = source.ClientSecret != ""
source.ClientSecret = ""
}
func GetAuthSources() ([]AuthSource, error) {
var sources []AuthSource
err := DB.Order("id asc").Find(&sources).Error
for index := range sources {
sources[index].Sanitize()
}
return sources, err
}
func GetActiveAuthSources() ([]AuthSource, error) {
var sources []AuthSource
err := DB.Where("is_active = ?", true).Order("id asc").Find(&sources).Error
for index := range sources {
sources[index].Sanitize()
}
return sources, err
}
func GetAuthSourceByID(id uint) (*AuthSource, error) {
if id == 0 {
return nil, errors.New("认证源 ID 不能为空")
}
var source AuthSource
if err := DB.First(&source, "id = ?", id).Error; err != nil {
return nil, err
}
source.ClientSecretConfigured = source.ClientSecret != ""
return &source, nil
}
func GetAuthSourceByName(name string) (*AuthSource, error) {
name = strings.TrimSpace(name)
if name == "" {
return nil, errors.New("认证源名称不能为空")
}
var source AuthSource
if err := DB.First(&source, "name = ?", name).Error; err != nil {
return nil, err
}
source.ClientSecretConfigured = source.ClientSecret != ""
return &source, nil
}
func CreateAuthSource(source *AuthSource) error {
if err := source.Validate(); err != nil {
return err
}
return DB.Create(source).Error
}
func UpdateAuthSource(source *AuthSource, keepSecret bool) error {
if source.ID == 0 {
return errors.New("认证源 ID 不能为空")
}
var current AuthSource
if err := DB.First(&current, "id = ?", source.ID).Error; err != nil {
return err
}
if keepSecret {
source.ClientSecret = current.ClientSecret
}
if err := source.Validate(); err != nil {
return err
}
return DB.Model(&current).Updates(map[string]any{
"name": source.Name,
"type": source.Type,
"display_name": source.DisplayName,
"is_active": source.IsActive,
"client_id": source.ClientID,
"client_secret": source.ClientSecret,
"openid_discovery_url": source.OpenIDDiscoveryURL,
"scopes": source.Scopes,
"icon_url": source.IconURL,
}).Error
}
func ToggleAuthSource(id uint, isActive bool) error {
source, err := GetAuthSourceByID(id)
if err != nil {
return err
}
source.IsActive = isActive
if err := source.Validate(); err != nil {
return err
}
return DB.Model(&AuthSource{}).Where("id = ?", id).Update("is_active", isActive).Error
}
func DeleteAuthSource(id uint) error {
if id == 0 {
return errors.New("认证源 ID 不能为空")
}
return DB.Transaction(func(tx *gorm.DB) error {
if err := tx.Where("auth_source_id = ?", id).Delete(&ExternalAccount{}).Error; err != nil {
return err
}
return tx.Delete(&AuthSource{}, "id = ?", id).Error
})
}
func FindExternalAccount(sourceID uint, externalID string) (*ExternalAccount, error) {
var account ExternalAccount
err := DB.Where("auth_source_id = ? AND external_id = ?", sourceID, externalID).First(&account).Error
if err != nil {
return nil, err
}
return &account, nil
}
func LinkExternalAccount(account *ExternalAccount) error {
if account.AuthSourceID == 0 || account.UserID == 0 || strings.TrimSpace(account.ExternalID) == "" {
return errors.New("外部账号绑定信息不完整")
}
account.ExternalID = strings.TrimSpace(account.ExternalID)
account.ExternalUsername = strings.TrimSpace(account.ExternalUsername)
account.Email = strings.TrimSpace(account.Email)
return DB.Where(ExternalAccount{
AuthSourceID: account.AuthSourceID,
ExternalID: account.ExternalID,
}).FirstOrCreate(account).Error
}
func ListExternalAccountsByUserID(userID int) ([]ExternalAccountView, error) {
if userID <= 0 {
return nil, errors.New("用户 ID 不能为空")
}
var accounts []ExternalAccount
if err := DB.Preload("AuthSource").Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil {
return nil, err
}
views := make([]ExternalAccountView, 0, len(accounts))
for _, account := range accounts {
label := account.AuthSource.DisplayName
if label == "" {
label = account.AuthSource.Name
}
views = append(views, ExternalAccountView{
ID: account.ID,
AuthSourceID: account.AuthSourceID,
AuthSourceName: account.AuthSource.Name,
AuthSourceType: account.AuthSource.Type,
AuthSourceLabel: label,
ExternalUsername: account.ExternalUsername,
Email: account.Email,
CreatedAt: account.CreatedAt,
})
}
return views, nil
}
func DeleteExternalAccountForUser(id uint, userID int) error {
if id == 0 {
return errors.New("绑定记录 ID 不能为空")
}
if userID <= 0 {
return errors.New("用户 ID 不能为空")
}
result := DB.Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{})
if result.Error != nil {
return result.Error
}
if result.RowsAffected == 0 {
return errors.New("绑定记录不存在")
}
return nil
}
@@ -0,0 +1,45 @@
package model
import "time"
type ConfigVersionSummary struct {
ID uint `json:"id"`
Version string `json:"version"`
Checksum string `json:"checksum"`
IsActive bool `json:"is_active"`
CreatedBy string `json:"created_by"`
CreatedAt time.Time `json:"created_at"`
}
type ConfigVersion struct {
ID uint `json:"id" gorm:"primaryKey"`
Version string `json:"version" gorm:"uniqueIndex;size:32;not null"`
SnapshotJSON string `json:"snapshot_json" gorm:"type:text;not null"`
MainConfig string `json:"main_config" gorm:"type:text;not null;default:''"`
RenderedConfig string `json:"rendered_config" gorm:"type:text;not null"`
SupportFilesJSON string `json:"support_files_json" gorm:"type:text;not null;default:'[]'"`
Checksum string `json:"checksum" gorm:"size:64;not null"`
IsActive bool `json:"is_active" gorm:"not null;default:false;index"`
CreatedBy string `json:"created_by" gorm:"size:64;not null"`
CreatedAt time.Time `json:"created_at"`
}
func ListConfigVersionSummaries() (versions []*ConfigVersionSummary, err error) {
err = DB.Model(&ConfigVersion{}).
Select("id", "version", "checksum", "is_active", "created_by", "created_at").
Order("id desc").
Find(&versions).Error
return versions, err
}
func GetConfigVersionByID(id uint) (*ConfigVersion, error) {
version := &ConfigVersion{}
err := DB.First(version, id).Error
return version, err
}
func GetActiveConfigVersion() (*ConfigVersion, error) {
version := &ConfigVersion{}
err := DB.Where("is_active = ?", true).Order("id desc").First(version).Error
return version, err
}
@@ -0,0 +1,27 @@
package model
import (
"time"
"github.com/rain-kl/openflare/openflare-server/internal/model/migrate"
)
const (
legacyDatabaseSchemaVersion = migrate.BaseDatabaseSchemaVersion
legacyMigrationTerminalVersion = 17
databaseSchemaVersionRowID = 1
)
// currentDatabaseSchemaVersion tracks the current physical schema validated by the
// legacy validator set. Goose owns only post-v17 migrations, and none exist yet.
var currentDatabaseSchemaVersion = legacyMigrationTerminalVersion
type DatabaseSchemaVersion struct {
ID uint `json:"id" gorm:"primaryKey"`
Version int `json:"version" gorm:"not null"`
UpdatedAt time.Time `json:"updated_at"`
}
func (DatabaseSchemaVersion) TableName() string {
return "database_schema_versions"
}
@@ -0,0 +1,35 @@
package model
import "time"
type DnsAccount struct {
ID uint `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"size:255;not null"`
Type string `json:"type" gorm:"size:64;not null"`
Authorization string `json:"-" gorm:"type:text;not null"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func ListDnsAccounts() (accounts []*DnsAccount, err error) {
err = DB.Order("id desc").Find(&accounts).Error
return accounts, err
}
func GetDnsAccountByID(id uint) (*DnsAccount, error) {
account := &DnsAccount{}
err := DB.First(account, id).Error
return account, err
}
func (account *DnsAccount) Insert() error {
return DB.Create(account).Error
}
func (account *DnsAccount) Update() error {
return DB.Save(account).Error
}
func (account *DnsAccount) Delete() error {
return DB.Delete(account).Error
}
@@ -0,0 +1,326 @@
package goose
import (
"database/sql"
"errors"
"fmt"
"io"
"log/slog"
"os"
"gorm.io/gorm"
)
type BridgeContext interface {
Context
AutoMigrateLegacySchemaMetadata(db *gorm.DB) error
InitializeFreshDatabaseSchema(db *gorm.DB, backend string) error
IsDatabaseEmpty(db *gorm.DB) (bool, error)
RepairCurrentSchemaState(db *gorm.DB, backend string) error
SaveLegacyDatabaseSchemaVersion(db *gorm.DB, version int) error
UpgradeLegacyDatabaseSchema(db *gorm.DB, backend string, version int) error
ValidateCurrentDatabaseSchema(db *gorm.DB, backend string) error
}
type schemaMigrationState int
const (
schemaMigrationStateFresh schemaMigrationState = iota
schemaMigrationStateLegacyOnly
schemaMigrationStateGooseOnly
schemaMigrationStateLegacyBootstrap
schemaMigrationStateMixed
)
func detectSchemaState(db *gorm.DB, ctx BridgeContext) (schemaMigrationState, error) {
hasLegacyTable := db.Migrator().HasTable("database_schema_versions")
hasGooseTable := db.Migrator().HasTable("goose_db_version")
switch {
case hasLegacyTable && hasGooseTable:
return schemaMigrationStateMixed, nil
case hasLegacyTable:
return schemaMigrationStateLegacyOnly, nil
case hasGooseTable:
return schemaMigrationStateGooseOnly, nil
}
empty, err := ctx.IsDatabaseEmpty(db)
if err != nil {
return 0, err
}
if empty {
return schemaMigrationStateFresh, nil
}
return schemaMigrationStateLegacyBootstrap, nil
}
func LoadDatabaseVersion(db *gorm.DB) (int, bool, error) {
if db == nil || !db.Migrator().HasTable("goose_db_version") {
return 0, false, nil
}
var version int64
err := db.Table("goose_db_version").
Where("is_applied = ?", true).
Order("version_id DESC").
Select("version_id").
Limit(1).
Row().
Scan(&version)
if errors.Is(err, sql.ErrNoRows) {
return 0, false, nil
}
if err != nil {
return 0, false, err
}
return int(version), true, nil
}
func loadLegacyDatabaseSchemaVersion(db *gorm.DB) (int, bool, error) {
if db == nil || !db.Migrator().HasTable("database_schema_versions") {
return 0, false, nil
}
var version int
err := db.Table("database_schema_versions").
Where("id = ?", 1).
Select("version").
Limit(1).
Row().
Scan(&version)
if errors.Is(err, sql.ErrNoRows) {
return 0, false, nil
}
if err != nil {
return 0, false, err
}
return version, true, nil
}
func bootstrapLegacySchemaVersion(db *gorm.DB, ctx BridgeContext) error {
if err := ctx.AutoMigrateLegacySchemaMetadata(db); err != nil {
return err
}
version, exists, err := loadLegacyDatabaseSchemaVersion(db)
if err != nil {
return err
}
if exists {
if int64(version) > LegacyBridgeVersion {
return fmt.Errorf("legacy schema version %d is newer than supported terminal version %d", version, LegacyBridgeVersion)
}
return nil
}
return ctx.SaveLegacyDatabaseSchemaVersion(db, 7)
}
func upgradeLegacyToTerminal(db *gorm.DB, backend string, ctx BridgeContext) error {
if err := bootstrapLegacySchemaVersion(db, ctx); err != nil {
return err
}
version, exists, err := loadLegacyDatabaseSchemaVersion(db)
if err != nil {
return err
}
if !exists {
return fmt.Errorf("legacy schema version record is missing after bootstrap")
}
return ctx.UpgradeLegacyDatabaseSchema(db, backend, version)
}
func validateGooseBridgeState(db *gorm.DB) error {
version, exists, err := LoadDatabaseVersion(db)
if err != nil {
return err
}
if !exists {
return nil
}
if int64(version) < LegacyBridgeVersion {
return fmt.Errorf("goose schema version %d is below legacy bridge baseline %d", version, LegacyBridgeVersion)
}
if int64(version) > CurrentTargetVersion() {
return fmt.Errorf("goose schema version %d is newer than application target version %d", version, CurrentTargetVersion())
}
return nil
}
func finalizeLegacyToGooseBridge(db *gorm.DB) error {
gooseVersion, exists, err := LoadDatabaseVersion(db)
if err != nil {
return err
}
if !exists || int64(gooseVersion) < LegacyBridgeVersion {
return nil
}
if !db.Migrator().HasTable("database_schema_versions") {
return nil
}
if err := db.Exec("DROP TABLE IF EXISTS database_schema_versions").Error; err != nil {
return fmt.Errorf("drop legacy schema versions table failed: %w", err)
}
slog.Info("completed legacy-to-goose migration bridge", "goose_version", gooseVersion)
return nil
}
func ValidateRegisteredSchema(db *gorm.DB) error {
if err := validateNodeCapabilitiesJSON(db); err != nil {
return err
}
return nil
}
func EnsureDatabaseSchemaUpToDate(db *gorm.DB, backend string, ctx BridgeContext) (returnedErr error) {
if backend == "sqlite" {
backupPath, restore, err := backupSQLiteDatabase(db)
if err != nil {
slog.Warn("failed to backup sqlite database before migration", "error", err)
} else if backupPath != "" {
defer func() {
if returnedErr != nil {
restore()
} else {
os.Remove(backupPath)
}
}()
}
}
var startDesc string
legacyVer, hasLegacy, _ := loadLegacyDatabaseSchemaVersion(db)
gooseVer, hasGoose, _ := LoadDatabaseVersion(db)
if hasGoose {
startDesc = fmt.Sprintf("goose version %d", gooseVer)
} else if hasLegacy {
startDesc = fmt.Sprintf("legacy version %d", legacyVer)
} else {
startDesc = "none (fresh database)"
}
state, err := detectSchemaState(db, ctx)
if err != nil {
return err
}
switch state {
case schemaMigrationStateFresh:
if err := ctx.InitializeFreshDatabaseSchema(db, backend); err != nil {
return err
}
case schemaMigrationStateLegacyOnly:
if err := upgradeLegacyToTerminal(db, backend, ctx); err != nil {
return err
}
case schemaMigrationStateGooseOnly:
if err := validateGooseBridgeState(db); err != nil {
return err
}
case schemaMigrationStateLegacyBootstrap:
if err := upgradeLegacyToTerminal(db, backend, ctx); err != nil {
return err
}
case schemaMigrationStateMixed:
legacyVersion, exists, err := loadLegacyDatabaseSchemaVersion(db)
if err != nil {
return err
}
if exists && int64(legacyVersion) != LegacyBridgeVersion {
return fmt.Errorf("incomplete mixed migration state: legacy schema version %d does not match bridge terminal version %d", legacyVersion, LegacyBridgeVersion)
}
if err := validateGooseBridgeState(db); err != nil {
return err
}
default:
return fmt.Errorf("unknown schema migration state: %d", state)
}
if err := runMigrations(db, backend, ctx); err != nil {
return err
}
if err := finalizeLegacyToGooseBridge(db); err != nil {
return err
}
if err := ctx.RepairCurrentSchemaState(db, backend); err != nil {
return err
}
if err := ctx.ValidateCurrentDatabaseSchema(db, backend); err != nil {
return err
}
if err := ValidateRegisteredSchema(db); err != nil {
return err
}
endVer, _, _ := LoadDatabaseVersion(db)
if hasGoose && int64(gooseVer) == int64(endVer) {
slog.Info("database schema is already up to date", "version", endVer)
} else {
slog.Info("database migration completed successfully", "from", startDesc, "to", fmt.Sprintf("goose version %d", endVer))
}
return nil
}
func backupSQLiteDatabase(db *gorm.DB) (string, func(), error) {
var dbList []struct {
Seq int
Name string
File string
}
if err := db.Raw("PRAGMA database_list").Scan(&dbList).Error; err != nil {
return "", nil, err
}
var dbPath string
for _, item := range dbList {
if item.Name == "main" && item.File != "" {
dbPath = item.File
break
}
}
if dbPath == "" {
return "", nil, nil
}
backupPath := dbPath + ".bak"
src, err := os.Open(dbPath)
if err != nil {
return "", nil, err
}
defer src.Close()
dst, err := os.Create(backupPath)
if err != nil {
return "", nil, err
}
defer dst.Close()
if _, err := io.Copy(dst, src); err != nil {
return "", nil, err
}
dst.Sync()
restoreFunc := func() {
src, err := os.Open(backupPath)
if err != nil {
slog.Error("failed to open sqlite backup for restore", "error", err)
return
}
defer src.Close()
dst, err := os.OpenFile(dbPath, os.O_WRONLY|os.O_TRUNC, 0644)
if err != nil {
slog.Error("failed to open sqlite db for restore", "error", err)
return
}
defer dst.Close()
if _, err := io.Copy(dst, src); err != nil {
slog.Error("failed to restore sqlite backup", "error", err)
} else {
dst.Sync()
slog.Warn("restored sqlite database from backup due to migration failure")
}
}
return backupPath, restoreFunc, nil
}
@@ -0,0 +1,50 @@
package goose
import (
"encoding/json"
"fmt"
presslygoose "github.com/pressly/goose/v3"
"gorm.io/gorm"
)
const versionNodeCapabilitiesJSON int64 = 202606020001
// migration202606020001 adds a future-proof JSON field for node capability
// summaries after the legacy v17 migration bridge.
func migration202606020001(backend string, ctx Context) *presslygoose.Migration {
return newGORMMigration(
versionNodeCapabilitiesJSON,
"202606020001_add_node_capabilities_json.go",
backend,
ctx,
migrateNodeCapabilitiesJSON,
)
}
func migrateNodeCapabilitiesJSON(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
emptyJSON, err := json.Marshal([]string{})
if err != nil {
return fmt.Errorf("marshal default node capabilities: %w", err)
}
if err := db.Exec(
`UPDATE nodes SET capabilities_json = ? WHERE capabilities_json IS NULL OR TRIM(capabilities_json) = ''`,
string(emptyJSON),
).Error; err != nil {
return fmt.Errorf("backfill nodes.capabilities_json: %w", err)
}
return validateNodeCapabilitiesJSON(db)
}
func validateNodeCapabilitiesJSON(db *gorm.DB) error {
if db == nil {
return fmt.Errorf("database handle is nil")
}
if !db.Migrator().HasColumn("nodes", "capabilities_json") {
return fmt.Errorf("column nodes.capabilities_json is missing")
}
return nil
}
@@ -0,0 +1,61 @@
package goose
import (
"fmt"
presslygoose "github.com/pressly/goose/v3"
"gorm.io/gorm"
)
const versionPagesStaticHosting int64 = 202606030001
// migration202606030001 adds OpenFlare Pages static hosting tables and the
// proxy_routes.pages_project_id binding used by the global release snapshot.
func migration202606030001(backend string, ctx Context) *presslygoose.Migration {
return newGORMMigration(
versionPagesStaticHosting,
"202606030001_add_pages_static_hosting.go",
backend,
ctx,
migratePagesStaticHosting,
)
}
func migratePagesStaticHosting(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
if err := db.Exec(
`UPDATE proxy_routes SET upstream_type = 'direct' WHERE upstream_type IS NULL OR TRIM(upstream_type) = ''`,
).Error; err != nil {
return fmt.Errorf("backfill proxy_routes.upstream_type: %w", err)
}
return validatePagesStaticHosting(db)
}
func validatePagesStaticHosting(db *gorm.DB) error {
if db == nil {
return fmt.Errorf("database handle is nil")
}
for _, table := range []string{"pages_projects", "pages_deployments", "pages_deployment_files"} {
if !db.Migrator().HasTable(table) {
return fmt.Errorf("table %s is missing", table)
}
}
for _, column := range []string{"upstream_type", "pages_project_id"} {
if !db.Migrator().HasColumn("proxy_routes", column) {
return fmt.Errorf("column proxy_routes.%s is missing", column)
}
}
for _, column := range []string{"slug", "active_deployment_id", "spa_fallback_enabled", "spa_fallback_path"} {
if !db.Migrator().HasColumn("pages_projects", column) {
return fmt.Errorf("column pages_projects.%s is missing", column)
}
}
for _, column := range []string{"project_id", "checksum", "artifact_path"} {
if !db.Migrator().HasColumn("pages_deployments", column) {
return fmt.Errorf("column pages_deployments.%s is missing", column)
}
}
return nil
}
@@ -0,0 +1,37 @@
package goose
import (
"fmt"
presslygoose "github.com/pressly/goose/v3"
"gorm.io/gorm"
)
const versionPagesSPAFallbackPath int64 = 202606030002
// migration202606030002 adds a configurable SPA fallback path for Pages
// projects. Existing projects keep the previous /index.html behavior.
func migration202606030002(backend string, ctx Context) *presslygoose.Migration {
return newGORMMigration(
versionPagesSPAFallbackPath,
"202606030002_add_pages_spa_fallback_path.go",
backend,
ctx,
migratePagesSPAFallbackPath,
)
}
func migratePagesSPAFallbackPath(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
if err := db.Exec(
`UPDATE pages_projects SET spa_fallback_path = '/index.html' WHERE spa_fallback_path IS NULL OR TRIM(spa_fallback_path) = ''`,
).Error; err != nil {
return fmt.Errorf("backfill pages_projects.spa_fallback_path: %w", err)
}
if !db.Migrator().HasColumn("pages_projects", "spa_fallback_path") {
return fmt.Errorf("column pages_projects.spa_fallback_path is missing")
}
return nil
}
@@ -0,0 +1,41 @@
package goose
import (
"fmt"
presslygoose "github.com/pressly/goose/v3"
"gorm.io/gorm"
)
const versionDropProxyRouteLegacyPoW int64 = 202606030003
// migration202606030003 drops the legacy pow_enabled and pow_config columns
// from proxy_routes table, since PoW is now entirely managed under WAF rule groups.
func migration202606030003(backend string, ctx Context) *presslygoose.Migration {
return newGORMMigration(
versionDropProxyRouteLegacyPoW,
"202606030003_drop_proxy_route_legacy_pow.go",
backend,
ctx,
migrateDropProxyRouteLegacyPoW,
)
}
func migrateDropProxyRouteLegacyPoW(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
// Drop pow_enabled column if exists
if db.Migrator().HasColumn("proxy_routes", "pow_enabled") {
if err := db.Exec("ALTER TABLE proxy_routes DROP COLUMN pow_enabled").Error; err != nil {
return fmt.Errorf("drop proxy_routes.pow_enabled: %w", err)
}
}
// Drop pow_config column if exists
if db.Migrator().HasColumn("proxy_routes", "pow_config") {
if err := db.Exec("ALTER TABLE proxy_routes DROP COLUMN pow_config").Error; err != nil {
return fmt.Errorf("drop proxy_routes.pow_config: %w", err)
}
}
return nil
}
@@ -0,0 +1,64 @@
package goose
import (
"fmt"
presslygoose "github.com/pressly/goose/v3"
"gorm.io/gorm"
)
const versionPagesFeaturesAndCleanup int64 = 202606040004
// migration202606040004 merges migrations 202606030004, 202606040001, 202606040002, and 202606040003.
// It adds Pages API proxying fields and RootDir/EntryFile to Pages projects,
// backfills default entry_file to 'index.html', and ensures unused fields (root_dir, entry_file)
// are dropped from Pages deployments.
func migration202606040004(backend string, ctx Context) *presslygoose.Migration {
return newGORMMigration(
versionPagesFeaturesAndCleanup,
"202606040004_add_pages_features_and_cleanup.go",
backend,
ctx,
migratePagesFeaturesAndCleanup,
)
}
func migratePagesFeaturesAndCleanup(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
// 1. Verify Pages projects columns
cols := []string{
"api_proxy_enabled", "api_proxy_path", "api_proxy_pass", "api_proxy_rewrite",
"root_dir", "entry_file",
}
for _, col := range cols {
if !db.Migrator().HasColumn("pages_projects", col) {
return fmt.Errorf("column pages_projects.%s is missing", col)
}
}
// 2. Backfill pages_projects.entry_file to 'index.html' if empty
type PagesProject struct {
ID uint `gorm:"primaryKey"`
EntryFile string `gorm:"size:512;not null;default:'index.html'"`
}
if err := db.Model(&PagesProject{}).Where("entry_file = '' OR entry_file IS NULL").Update("entry_file", "index.html").Error; err != nil {
return fmt.Errorf("failed to backfill pages_projects.entry_file: %w", err)
}
// 3. Drop unused fields root_dir and entry_file from pages_deployments if they exist
if db.Migrator().HasColumn("pages_deployments", "root_dir") {
if err := db.Exec("ALTER TABLE pages_deployments DROP COLUMN root_dir").Error; err != nil {
return fmt.Errorf("failed to drop pages_deployments.root_dir: %w", err)
}
}
if db.Migrator().HasColumn("pages_deployments", "entry_file") {
if err := db.Exec("ALTER TABLE pages_deployments DROP COLUMN entry_file").Error; err != nil {
return fmt.Errorf("failed to drop pages_deployments.entry_file: %w", err)
}
}
return nil
}
@@ -0,0 +1,75 @@
package goose
import (
"context"
"database/sql"
"fmt"
presslygoose "github.com/pressly/goose/v3"
"gorm.io/gorm"
)
const LegacyBridgeVersion int64 = 17
type migrationFunc func(ctx Context, db *gorm.DB, backend string) error
func newBaselineMigration() *presslygoose.Migration {
migration := presslygoose.NewGoMigration(LegacyBridgeVersion, nil, nil)
migration.Source = fmt.Sprintf("%05d_legacy_terminal_baseline.go", LegacyBridgeVersion)
return migration
}
func newGORMMigration(version int64, source string, backend string, ctx Context, up migrationFunc) *presslygoose.Migration {
migration := presslygoose.NewGoMigration(version, &presslygoose.GoFunc{
RunDB: func(_ context.Context, sqlDB *sql.DB) error {
gormDB, err := openGORMDB(ctx, sqlDB, backend)
if err != nil {
return err
}
if backend == "postgres" {
return gormDB.Transaction(func(tx *gorm.DB) error {
return up(ctx, tx, backend)
})
}
return up(ctx, gormDB, backend)
},
}, nil)
migration.Source = source
return migration
}
func registeredMigrations(backend string, ctx Context) []*presslygoose.Migration {
return []*presslygoose.Migration{
migration202606020001(backend, ctx),
migration202606030001(backend, ctx),
migration202606030002(backend, ctx),
migration202606030003(backend, ctx),
migration202606040004(backend, ctx),
}
}
func buildMigrations(backend string, ctx Context) []*presslygoose.Migration {
migrations := []*presslygoose.Migration{newBaselineMigration()}
migrations = append(migrations, registeredMigrations(backend, ctx)...)
return migrations
}
func CurrentTargetVersion() int64 {
var maxVersion int64 = LegacyBridgeVersion
for _, migration := range buildMigrations("sqlite", noopContext{}) {
if migration.Version > maxVersion {
maxVersion = migration.Version
}
}
return maxVersion
}
type noopContext struct{}
func (noopContext) ApplyCurrentSchema(db *gorm.DB, backend string) error {
return nil
}
func (noopContext) RegisterSharding(db *gorm.DB, backend string) error {
return nil
}
@@ -0,0 +1,81 @@
package goose
import (
"context"
"database/sql"
"fmt"
"github.com/glebarez/sqlite"
presslygoose "github.com/pressly/goose/v3"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/schema"
)
type Context interface {
ApplyCurrentSchema(db *gorm.DB, backend string) error
RegisterSharding(db *gorm.DB, backend string) error
}
func dialectForBackend(backend string) (presslygoose.Dialect, error) {
switch backend {
case "postgres":
return presslygoose.DialectPostgres, nil
case "sqlite":
return presslygoose.DialectSQLite3, nil
default:
return "", fmt.Errorf("unsupported database backend: %s", backend)
}
}
func openGORMDB(ctx Context, db *sql.DB, backend string) (*gorm.DB, error) {
var dialector gorm.Dialector
switch backend {
case "postgres":
dialector = postgres.New(postgres.Config{Conn: db})
case "sqlite":
dialector = &sqlite.Dialector{Conn: db}
default:
return nil, fmt.Errorf("unsupported database backend: %s", backend)
}
gormDB, err := gorm.Open(dialector, &gorm.Config{
NamingStrategy: schema.NamingStrategy{},
})
if err != nil {
return nil, err
}
if err := ctx.RegisterSharding(gormDB, backend); err != nil {
return nil, err
}
return gormDB, nil
}
func buildProvider(db *gorm.DB, backend string, ctx Context) (*presslygoose.Provider, error) {
sqlDB, err := db.DB()
if err != nil {
return nil, err
}
dialect, err := dialectForBackend(backend)
if err != nil {
return nil, err
}
return presslygoose.NewProvider(
dialect,
sqlDB,
nil,
presslygoose.WithDisableGlobalRegistry(true),
presslygoose.WithGoMigrations(buildMigrations(backend, ctx)...),
)
}
func runMigrations(db *gorm.DB, backend string, ctx Context) error {
provider, err := buildProvider(db, backend, ctx)
if err != nil {
return fmt.Errorf("build goose provider: %w", err)
}
if _, err := provider.Up(context.Background()); err != nil {
return fmt.Errorf("goose up failed: %w", err)
}
return nil
}
+369
View File
@@ -0,0 +1,369 @@
package model
import (
"fmt"
"log/slog"
"os"
"reflect"
"strings"
"sync"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/utils/security"
"github.com/glebarez/sqlite"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/schema"
)
var DB *gorm.DB
type dbModel struct {
value any
tableName string
hasIDPK bool
}
func registeredModels() []any {
return []any{
&User{},
&AuthSource{},
&ExternalAccount{},
&Option{},
&Origin{},
&ProxyRoute{},
&PagesProject{},
&PagesDeployment{},
&PagesDeploymentFile{},
&ConfigVersion{},
&Node{},
&NodeSystemProfile{},
&ApplyLog{},
&NodeMetricSnapshot{},
&NodeRequestReport{},
&NodeAccessLog{},
&NodeHealthEvent{},
&NodeObservationOpenresty{},
&NodeObservationFrps{},
&NodeObservationFrpc{},
&TLSCertificate{},
&ManagedDomain{},
&AcmeAccount{},
&DnsAccount{},
&WAFRuleGroup{},
&WAFIPGroup{},
&WAFRuleGroupBinding{},
}
}
func currentSchemaMetadataModels() []any {
return nil
}
func legacySchemaMetadataModels() []any {
return []any{
&DatabaseSchemaVersion{},
}
}
func schemaMetadataModels() []any {
models := make([]any, 0, len(currentSchemaMetadataModels())+len(legacySchemaMetadataModels()))
models = append(models, currentSchemaMetadataModels()...)
models = append(models, legacySchemaMetadataModels()...)
return models
}
func buildDBModels() ([]dbModel, error) {
models := registeredModels()
result := make([]dbModel, 0, len(models))
namer := schema.NamingStrategy{}
cache := &sync.Map{}
for _, item := range models {
parsed, err := schema.Parse(item, cache, namer)
if err != nil {
return nil, err
}
hasIDPK := len(parsed.PrimaryFields) == 1 && parsed.PrimaryFields[0].DBName == "id"
result = append(result, dbModel{
value: item,
tableName: parsed.Table,
hasIDPK: hasIDPK,
})
}
return result, nil
}
func createRootAccountIfNeed() error {
var user User
//if user.Status != common.UserStatusEnabled {
if err := DB.First(&user).Error; err != nil {
slog.Info("no user exists, create a root user", "username", "root")
hashedPassword, err := security.Password2Hash("123456")
if err != nil {
return err
}
rootUser := User{
Username: "root",
Password: hashedPassword,
Role: common.RoleRootUser,
Status: common.UserStatusEnabled,
DisplayName: "Root User",
}
DB.Create(&rootUser)
}
return nil
}
func CountTable(tableName string) (num int64) {
DB.Table(tableName).Count(&num)
return
}
func openDatabase() (*gorm.DB, string, error) {
if common.SQLDSN != "" {
db, err := gorm.Open(postgres.Open(common.SQLDSN), &gorm.Config{})
if err != nil {
return nil, "", err
}
return db, "postgres", nil
}
db, err := gorm.Open(sqlite.Open(common.SQLitePath), &gorm.Config{})
if err != nil {
return nil, "", err
}
slog.Info("database DSN not set, using SQLite as database", "sqlite_path", common.SQLitePath)
return db, "sqlite", nil
}
func autoMigrateAll(db *gorm.DB) error {
return autoMigrateAllExcept(db, nil)
}
func autoMigrateAllExcept(db *gorm.DB, excludedTables map[string]bool) error {
models := registeredModels()
for i, item := range models {
name := fmt.Sprintf("%T", item)
tableName, err := tableNameForModel(item)
if err != nil {
return fmt.Errorf("resolve table name for %s failed: %w", name, err)
}
if excludedTables[tableName] {
slog.Info("autoMigrateAll: skipped model", "index", fmt.Sprintf("%d/%d", i+1, len(models)), "model", name, "table", tableName)
continue
}
slog.Info("autoMigrateAll: migrating model", "index", fmt.Sprintf("%d/%d", i+1, len(models)), "model", name)
if err := db.AutoMigrate(item); err != nil {
return fmt.Errorf("AutoMigrate %s failed: %w", name, err)
}
slog.Info("autoMigrateAll: migrated model", "model", name)
}
return nil
}
func tableNameForModel(item any) (string, error) {
namer := schema.NamingStrategy{}
cache := &sync.Map{}
parsed, err := schema.Parse(item, cache, namer)
if err != nil {
return "", err
}
return parsed.Table, nil
}
func isDatabaseEmpty(db *gorm.DB) (bool, error) {
models, err := buildDBModels()
if err != nil {
return false, err
}
for _, item := range models {
if isShardedObservabilityTable(item.tableName) {
for _, table := range observabilityShardTables(item.tableName) {
if !db.Migrator().HasTable(table) {
continue
}
var count int64
if err := db.Table(table).Limit(1).Count(&count).Error; err != nil {
return false, err
}
if count > 0 {
return false, nil
}
}
continue
}
if !db.Migrator().HasTable(item.value) {
continue
}
var count int64
if err := db.Model(item.value).Limit(1).Count(&count).Error; err != nil {
return false, err
}
if count > 0 {
return false, nil
}
}
return true, nil
}
func sqliteSourceExists() bool {
info, err := os.Stat(common.SQLitePath)
if err != nil {
return false
}
return !info.IsDir()
}
func migrateSQLiteDataIfNeeded(target *gorm.DB, backend string) error {
if backend != "postgres" {
return nil
}
empty, err := isDatabaseEmpty(target)
if err != nil {
return err
}
if !empty {
slog.Info("skip sqlite migration because target database already has data", "backend", backend)
return nil
}
if !sqliteSourceExists() {
slog.Info("skip sqlite migration because sqlite source file was not found", "sqlite_path", common.SQLitePath)
return nil
}
source, err := gorm.Open(sqlite.Open(common.SQLitePath), &gorm.Config{
PrepareStmt: true,
})
if err != nil {
return fmt.Errorf("open sqlite source database failed: %w", err)
}
sourceSQLDB, err := source.DB()
if err != nil {
return fmt.Errorf("get sqlite source database handle failed: %w", err)
}
defer func() {
_ = sourceSQLDB.Close()
}()
models, err := buildDBModels()
if err != nil {
return err
}
slog.Info("starting sqlite to postgres database migration", "sqlite_path", common.SQLitePath)
err = target.Transaction(func(tx *gorm.DB) error {
for _, item := range models {
if err := migrateTableData(source, tx, item); err != nil {
return err
}
if item.hasIDPK {
if err := resetPostgresSequence(tx, item.tableName); err != nil {
return err
}
}
}
return nil
})
if err != nil {
return err
}
slog.Info("sqlite to postgres database migration completed", "sqlite_path", common.SQLitePath)
return nil
}
func migrateTableData(source *gorm.DB, target *gorm.DB, item dbModel) error {
if !source.Migrator().HasTable(item.value) {
slog.Info("database migration progress", "table", item.tableName, "migrated", 0, "total", 0, "status", "skipped_missing_source_table")
return nil
}
var total int64
if err := source.Model(item.value).Count(&total).Error; err != nil {
return fmt.Errorf("count sqlite table %s failed: %w", item.tableName, err)
}
slog.Info("database migration progress", "table", item.tableName, "migrated", 0, "total", total, "status", "starting")
if total == 0 {
slog.Info("database migration progress", "table", item.tableName, "migrated", 0, "total", total, "status", "completed")
return nil
}
modelType := reflect.TypeOf(item.value).Elem()
sliceType := reflect.SliceOf(modelType)
migrated := int64(0)
offset := 0
const batchSize = 200
for {
batchPtr := reflect.New(sliceType)
query := source.Model(item.value).Limit(batchSize).Offset(offset)
if item.hasIDPK {
query = query.Order("id ASC")
}
if err := query.Find(batchPtr.Interface()).Error; err != nil {
return fmt.Errorf("read sqlite table %s failed: %w", item.tableName, err)
}
batchLen := batchPtr.Elem().Len()
if batchLen == 0 {
break
}
if isShardedObservabilityTable(item.tableName) {
for index := 0; index < batchLen; index++ {
record := batchPtr.Elem().Index(index)
if err := target.Create(record.Addr().Interface()).Error; err != nil {
return fmt.Errorf("write target sharded table %s failed: %w", item.tableName, err)
}
}
} else {
if err := target.Create(batchPtr.Interface()).Error; err != nil {
return fmt.Errorf("write target table %s failed: %w", item.tableName, err)
}
}
migrated += int64(batchLen)
offset += batchLen
slog.Info("database migration progress", "table", item.tableName, "migrated", migrated, "total", total, "status", "running")
}
slog.Info("database migration progress", "table", item.tableName, "migrated", migrated, "total", total, "status", "completed")
return nil
}
func resetPostgresSequence(db *gorm.DB, tableName string) error {
sql := fmt.Sprintf(
"SELECT setval(pg_get_serial_sequence('%s', 'id'), COALESCE(MAX(id), 1), MAX(id) IS NOT NULL) FROM \"%s\"",
tableName,
tableName,
)
return db.Exec(sql).Error
}
func InitDB() (err error) {
db, backend, err := openDatabase()
if err != nil {
slog.Error("open database failed", "error", err)
os.Exit(1)
}
DB = db
if err = registerSharding(db, backend); err != nil {
return err
}
if err = ensureDatabaseSchemaUpToDate(db, backend); err != nil {
return err
}
return createRootAccountIfNeed()
}
func CloseDB() error {
sqlDB, err := DB.DB()
if err != nil {
return err
}
err = sqlDB.Close()
return err
}
func IsUniqueConstraintError(err error) bool {
if err == nil {
return false
}
return strings.Contains(strings.ToLower(err.Error()), "unique")
}
@@ -0,0 +1,886 @@
package model
import (
"encoding/json"
"go/ast"
"go/parser"
"go/token"
"os"
"path/filepath"
"reflect"
"strings"
"testing"
"time"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
)
type legacyProxyRouteV7 struct {
ID uint `gorm:"primaryKey"`
SiteName string `gorm:"size:255;not null;default:''"`
Domain string `gorm:"uniqueIndex;size:255;not null"`
Domains string `gorm:"type:text;not null;default:'[]'"`
OriginID *uint `gorm:"index"`
OriginURL string `gorm:"size:2048;not null"`
OriginHost string `gorm:"size:255"`
Upstreams string `gorm:"type:text;not null;default:'[]'"`
Enabled bool `gorm:"not null;default:true"`
EnableHTTPS bool `gorm:"column:enable_https;not null;default:false"`
CertID *uint
CertIDs string `gorm:"type:text;not null;default:'[]'"`
RedirectHTTP bool `gorm:"not null;default:false"`
LimitConnPerServer int `gorm:"not null;default:0"`
LimitConnPerIP int `gorm:"not null;default:0"`
LimitRate string `gorm:"size:32;not null;default:''"`
CacheEnabled bool `gorm:"not null;default:false"`
CachePolicy string `gorm:"size:32;not null;default:''"`
CacheRules string `gorm:"type:text;not null;default:'[]'"`
CustomHeaders string `gorm:"type:text;not null;default:'[]'"`
Remark string `gorm:"size:255"`
CreatedAt time.Time
UpdatedAt time.Time
}
func (legacyProxyRouteV7) TableName() string {
return "proxy_routes"
}
func openBareTestSQLiteDB(t *testing.T, name string) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), name)), &gorm.Config{})
if err != nil {
t.Fatalf("open sqlite db: %v", err)
}
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("get sql db: %v", err)
}
t.Cleanup(func() {
_ = sqlDB.Close()
})
return db
}
func openTestSQLiteDB(t *testing.T, name string) *gorm.DB {
t.Helper()
db := openBareTestSQLiteDB(t, name)
if err := autoMigrateAll(db); err != nil {
t.Fatalf("auto migrate db: %v", err)
}
return db
}
func findDBModelByTableName(t *testing.T, tableName string) dbModel {
t.Helper()
models, err := buildDBModels()
if err != nil {
t.Fatalf("build db models: %v", err)
}
for _, item := range models {
if item.tableName == tableName {
return item
}
}
t.Fatalf("db model not found for table %s", tableName)
return dbModel{}
}
func expectedCurrentDatabaseVersion() int {
return int(currentGooseTargetVersion())
}
func TestIsDatabaseEmpty(t *testing.T) {
db := openTestSQLiteDB(t, "empty.db")
empty, err := isDatabaseEmpty(db)
if err != nil {
t.Fatalf("isDatabaseEmpty returned error: %v", err)
}
if !empty {
t.Fatal("expected database to be empty")
}
if err := db.Create(&User{
Username: "alice",
Password: "secret",
DisplayName: "Alice",
Role: 1,
Status: 1,
}).Error; err != nil {
t.Fatalf("seed user: %v", err)
}
empty, err = isDatabaseEmpty(db)
if err != nil {
t.Fatalf("isDatabaseEmpty after seed returned error: %v", err)
}
if empty {
t.Fatal("expected database to be non-empty")
}
}
func TestMigrateTableDataCopiesRows(t *testing.T) {
source := openTestSQLiteDB(t, "source.db")
target := openTestSQLiteDB(t, "target.db")
user := User{
Id: 1,
Username: "root",
Password: "hashed",
DisplayName: "Root User",
Role: 100,
Status: 1,
}
option := Option{
Key: "AgentHeartbeatInterval",
Value: "10000",
}
if err := source.Create(&user).Error; err != nil {
t.Fatalf("seed source user: %v", err)
}
if err := source.Create(&option).Error; err != nil {
t.Fatalf("seed source option: %v", err)
}
if err := migrateTableData(source, target, findDBModelByTableName(t, "users")); err != nil {
t.Fatalf("migrate users: %v", err)
}
if err := migrateTableData(source, target, findDBModelByTableName(t, "options")); err != nil {
t.Fatalf("migrate options: %v", err)
}
var gotUser User
if err := target.First(&gotUser, 1).Error; err != nil {
t.Fatalf("query migrated user: %v", err)
}
if gotUser.Username != user.Username || gotUser.DisplayName != user.DisplayName {
t.Fatalf("unexpected migrated user: %+v", gotUser)
}
var gotOption Option
if err := target.First(&gotOption, "key = ?", option.Key).Error; err != nil {
t.Fatalf("query migrated option: %v", err)
}
if gotOption.Value != option.Value {
t.Fatalf("unexpected migrated option value: %s", gotOption.Value)
}
}
func TestRegisterShardingAutoMigratesShardTables(t *testing.T) {
db := openBareTestSQLiteDB(t, "sharded.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := autoMigrateAll(db); err != nil {
t.Fatalf("auto migrate db: %v", err)
}
for _, table := range []string{
"node_metric_snapshots_00",
"node_metric_snapshots_09",
"node_request_reports_00",
"node_request_reports_09",
"node_access_logs_00",
"node_access_logs_09",
} {
if !db.Migrator().HasTable(table) {
t.Fatalf("expected sharded table %s to exist", table)
}
}
}
func TestUpgradeDatabaseSchemaV15ToV16AppliesCompressedReleaseSchema(t *testing.T) {
db := openBareTestSQLiteDB(t, "v16.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := autoMigrateLegacySchemaMetadata(db); err != nil {
t.Fatalf("auto migrate legacy schema metadata: %v", err)
}
if err := applyCurrentSchema(db, "sqlite"); err != nil {
t.Fatalf("apply current schema: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
t.Fatalf("failed to add legacy pow_enabled: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
t.Fatalf("failed to add legacy pow_config: %v", err)
}
if err := ensureDefaultWAFRuleGroup(db); err != nil {
t.Fatalf("ensure default waf rule group: %v", err)
}
if err := saveDatabaseSchemaVersion(db, 15); err != nil {
t.Fatalf("save schema version: %v", err)
}
if err := upgradeDatabaseSchema(db, "sqlite", 15); err != nil {
t.Fatalf("upgrade schema: %v", err)
}
if !db.Migrator().HasTable(&WAFIPGroup{}) {
t.Fatal("expected waf_ip_groups table")
}
if !db.Migrator().HasColumn(&WAFRuleGroup{}, "ip_whitelist_groups") {
t.Fatal("expected waf_rule_groups.ip_whitelist_groups column")
}
if !db.Migrator().HasColumn(&Node{}, "access_token") {
t.Fatal("expected nodes.access_token column")
}
if !db.Migrator().HasColumn(&Node{}, "version") {
t.Fatal("expected nodes.version column")
}
if !db.Migrator().HasColumn(&Node{}, "ext_version") {
t.Fatal("expected nodes.ext_version column")
}
if !db.Migrator().HasColumn(&ProxyRoute{}, "tunnel_node_id") {
t.Fatal("expected proxy_routes.tunnel_node_id column")
}
if db.Migrator().HasTable("tunnels") {
t.Fatal("expected pre-release tunnels table to be absent")
}
version, ok, err := loadDatabaseSchemaVersion(db)
if err != nil {
t.Fatalf("load schema version: %v", err)
}
if !ok || version != currentDatabaseSchemaVersion {
t.Fatalf("unexpected schema version: got %d ok=%v want %d", version, ok, currentDatabaseSchemaVersion)
}
}
func TestMigrateObservabilityLegacyColumnsBackfillsHealthEventMetadata(t *testing.T) {
db := openTestSQLiteDB(t, "legacy-health-events.db")
if err := db.Exec("ALTER TABLE node_health_events ADD COLUMN raw_json TEXT").Error; err != nil {
t.Fatalf("add raw_json column: %v", err)
}
rawJSON, err := json.Marshal(map[string]any{
"event_type": "sync_error",
"metadata": map[string]string{
"reason": "checksum_mismatch",
"scope": "routes",
},
})
if err != nil {
t.Fatalf("marshal raw json: %v", err)
}
event := &NodeHealthEvent{
NodeID: "node-legacy",
EventType: "sync_error",
Severity: "warning",
Status: "active",
Message: "checksum mismatch",
FirstTriggeredAt: time.Now().Add(-time.Minute),
LastTriggeredAt: time.Now(),
ReportedAt: time.Now(),
}
if err := db.Create(event).Error; err != nil {
t.Fatalf("create health event: %v", err)
}
if err := db.Exec("UPDATE node_health_events SET raw_json = ? WHERE id = ?", string(rawJSON), event.ID).Error; err != nil {
t.Fatalf("seed legacy raw_json: %v", err)
}
if err := migrateObservabilityLegacyColumns(db); err != nil {
t.Fatalf("migrateObservabilityLegacyColumns: %v", err)
}
var got NodeHealthEvent
if err := db.First(&got, event.ID).Error; err != nil {
t.Fatalf("query health event: %v", err)
}
if got.MetadataJSON == "" {
t.Fatal("expected metadata_json to be backfilled")
}
}
func TestEnsureDatabaseSchemaUpToDateInitializesFreshDatabase(t *testing.T) {
db := openBareTestSQLiteDB(t, "fresh-schema.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
}
version, exists, err := loadDatabaseSchemaVersion(db)
if err != nil {
t.Fatalf("loadDatabaseSchemaVersion: %v", err)
}
if !exists {
t.Fatal("expected database schema version to be recorded")
}
if version != expectedCurrentDatabaseVersion() {
t.Fatalf("unexpected schema version: got %d want %d", version, expectedCurrentDatabaseVersion())
}
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
t.Fatal("expected fresh database to avoid legacy database_schema_versions table")
}
if !db.Migrator().HasTable("goose_db_version") {
t.Fatal("expected fresh database to initialize goose_db_version")
}
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
t.Fatal("expected fresh database to apply goose migration nodes.capabilities_json")
}
}
func TestEnsureDatabaseSchemaUpToDateUpgradesLegacyDatabase(t *testing.T) {
db := openBareTestSQLiteDB(t, "legacy-schema.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := autoMigrateAll(db); err != nil {
t.Fatalf("auto migrate db: %v", err)
}
// Add legacy PoW columns manually to proxy_routes table to simulate legacy schema v9-v17 state
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
t.Fatalf("failed to add legacy pow_enabled: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
t.Fatalf("failed to add legacy pow_config: %v", err)
}
if err := db.Create(&User{
Username: "legacy",
Password: "secret",
DisplayName: "Legacy User",
Role: 1,
Status: 1,
}).Error; err != nil {
t.Fatalf("seed legacy user: %v", err)
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
}
version, exists, err := loadDatabaseSchemaVersion(db)
if err != nil {
t.Fatalf("loadDatabaseSchemaVersion: %v", err)
}
if !exists {
t.Fatal("expected legacy database to gain a schema version record")
}
if version != expectedCurrentDatabaseVersion() {
t.Fatalf("unexpected schema version: got %d want %d", version, expectedCurrentDatabaseVersion())
}
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
t.Fatal("expected legacy database_schema_versions table to be removed after bridging to goose")
}
if !db.Migrator().HasTable("goose_db_version") {
t.Fatal("expected legacy upgrade to initialize goose_db_version")
}
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
t.Fatal("expected legacy upgrade to apply goose migration nodes.capabilities_json")
}
}
func TestMigrateOriginsSchemaBackfillsOrigins(t *testing.T) {
db := openBareTestSQLiteDB(t, "legacy-origins.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := applyCurrentSchema(db, "sqlite"); err != nil {
t.Fatalf("applyCurrentSchema: %v", err)
}
now := time.Now().UTC()
route := &ProxyRoute{
Domain: "app.example.com",
OriginURL: "https://origin-a.internal:8443/api",
Upstreams: `["https://origin-a.internal:8443/api"]`,
Enabled: true,
CreatedAt: now,
UpdatedAt: now,
}
if err := db.Create(route).Error; err != nil {
t.Fatalf("seed proxy route: %v", err)
}
if err := db.Exec(`DELETE FROM origins`).Error; err != nil {
t.Fatalf("clear origins: %v", err)
}
if err := db.Model(&ProxyRoute{}).Where("id = ?", route.ID).Update("origin_id", nil).Error; err != nil {
t.Fatalf("clear route origin_id: %v", err)
}
if err := backfillOriginsFromProxyRoutes(db); err != nil {
t.Fatalf("backfillOriginsFromProxyRoutes: %v", err)
}
if !db.Migrator().HasTable(&Origin{}) {
t.Fatal("expected origins table to exist")
}
if !db.Migrator().HasColumn(&ProxyRoute{}, "origin_id") {
t.Fatal("expected proxy_routes.origin_id column to exist")
}
reloadedRoute := &ProxyRoute{}
if err := db.First(reloadedRoute, route.ID).Error; err != nil {
t.Fatalf("query proxy route: %v", err)
}
if reloadedRoute.OriginID == nil || *reloadedRoute.OriginID == 0 {
t.Fatal("expected migrated route to be linked to a backfilled origin")
}
origin := &Origin{}
if err := db.First(origin, *reloadedRoute.OriginID).Error; err != nil {
t.Fatalf("query origin: %v", err)
}
if origin.Address != "origin-a.internal" {
t.Fatalf("unexpected backfilled origin address: %s", origin.Address)
}
}
func TestEnsureDatabaseSchemaUpToDateAddsProxyRouteDomainCertificateFields(t *testing.T) {
db := openBareTestSQLiteDB(t, "legacy-proxy-route-domain-cert-ids.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := autoMigrateLegacySchemaMetadata(db); err != nil {
t.Fatalf("auto migrate legacy schema metadata: %v", err)
}
for _, item := range registeredModels() {
if _, ok := item.(*ProxyRoute); ok {
continue
}
if err := db.AutoMigrate(item); err != nil {
t.Fatalf("auto migrate supporting table: %v", err)
}
}
if err := db.AutoMigrate(&legacyProxyRouteV7{}); err != nil {
t.Fatalf("auto migrate legacy proxy_routes v7: %v", err)
}
// Add legacy PoW columns manually to proxy_routes table to simulate legacy schema v9-v17 state
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
t.Fatalf("failed to add legacy pow_enabled: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
t.Fatalf("failed to add legacy pow_config: %v", err)
}
now := time.Now().UTC()
certID := uint(9)
if err := db.Create(&legacyProxyRouteV7{
SiteName: "secure-site",
Domain: "secure.example.com",
Domains: `["secure.example.com","www.secure.example.com"]`,
OriginURL: "https://origin-secure.internal:8443",
Upstreams: `["https://origin-secure.internal:8443"]`,
Enabled: true,
EnableHTTPS: true,
CertID: &certID,
CertIDs: `[9]`,
RedirectHTTP: true,
LimitConnPerServer: 120,
LimitConnPerIP: 12,
LimitRate: "512k",
CacheEnabled: false,
CachePolicy: "",
CacheRules: `[]`,
CustomHeaders: `[]`,
CreatedAt: now,
UpdatedAt: now,
}).Error; err != nil {
t.Fatalf("seed legacy proxy route v7: %v", err)
}
if err := saveDatabaseSchemaVersion(db, 7); err != nil {
t.Fatalf("save schema version: %v", err)
}
previousDB := DB
DB = db
t.Cleanup(func() {
DB = previousDB
})
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
}
var route ProxyRoute
if err := db.First(&route).Error; err != nil {
t.Fatalf("query migrated proxy route: %v", err)
}
var domainCertIDs []uint
if err := json.Unmarshal([]byte(route.DomainCertIDs), &domainCertIDs); err != nil {
t.Fatalf("decode migrated domain_cert_ids: %v", err)
}
if len(domainCertIDs) != 2 || domainCertIDs[0] != certID || domainCertIDs[1] != certID {
t.Fatalf("unexpected migrated domain_cert_ids: %#v", domainCertIDs)
}
}
func TestRunDatabaseSchemaMigrationDoesNotAdvanceVersionWhenValidationFails(t *testing.T) {
db := openBareTestSQLiteDB(t, "failed-validation.db")
err := runDatabaseSchemaMigration(db, "sqlite", databaseSchemaMigration{
fromVersion: legacyDatabaseSchemaVersion,
toVersion: 11,
migrate: func(tx *gorm.DB, backend string) error {
return autoMigrateLegacySchemaMetadata(tx)
},
validate: func(tx *gorm.DB, backend string) error {
return gorm.ErrInvalidDB
},
})
if err == nil {
t.Fatal("expected migration validation to fail")
}
_, exists, loadErr := loadDatabaseSchemaVersion(db)
if loadErr != nil {
t.Fatalf("loadDatabaseSchemaVersion: %v", loadErr)
}
if exists {
t.Fatal("expected schema version to remain unset after failed validation")
}
}
func TestEnsureDatabaseSchemaUpToDateAddsNodeIPManualOverride(t *testing.T) {
db := openBareTestSQLiteDB(t, "node-ip-manual-override-migration.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := applyCurrentSchema(db, "sqlite"); err != nil {
t.Fatalf("apply current schema: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
t.Fatalf("failed to add legacy pow_enabled: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
t.Fatalf("failed to add legacy pow_config: %v", err)
}
if err := ensureDefaultWAFRuleGroup(db); err != nil {
t.Fatalf("ensure default waf rule group: %v", err)
}
if err := db.Migrator().DropColumn(&Node{}, "ip_manual_override"); err != nil {
t.Fatalf("drop ip_manual_override column: %v", err)
}
if db.Migrator().HasColumn(&Node{}, "ip_manual_override") {
t.Fatal("expected test database to simulate schema v14 without ip_manual_override")
}
if err := saveDatabaseSchemaVersion(db, 14); err != nil {
t.Fatalf("save schema version: %v", err)
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
}
if !db.Migrator().HasColumn(&Node{}, "ip_manual_override") {
t.Fatal("expected migration to add nodes.ip_manual_override")
}
version, exists, err := loadDatabaseSchemaVersion(db)
if err != nil {
t.Fatalf("loadDatabaseSchemaVersion: %v", err)
}
if !exists {
t.Fatal("expected schema version record to exist")
}
if version != expectedCurrentDatabaseVersion() {
t.Fatalf("unexpected schema version: got %d want %d", version, expectedCurrentDatabaseVersion())
}
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
t.Fatal("expected migration chain to include nodes.capabilities_json")
}
}
func TestEnsureDatabaseSchemaUpToDateV16BackfillsNodeColumnsWhenNewColumnsAlreadyExist(t *testing.T) {
db := openBareTestSQLiteDB(t, "node-v16-existing-target-columns.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := applyCurrentSchema(db, "sqlite"); err != nil {
t.Fatalf("apply current schema: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
t.Fatalf("failed to add legacy pow_enabled: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
t.Fatalf("failed to add legacy pow_config: %v", err)
}
if err := ensureDefaultWAFRuleGroup(db); err != nil {
t.Fatalf("ensure default waf rule group: %v", err)
}
for _, stmt := range []string{
`ALTER TABLE nodes ADD COLUMN agent_token text`,
`ALTER TABLE nodes ADD COLUMN agent_version text`,
`ALTER TABLE nodes ADD COLUMN nginx_version text`,
`ALTER TABLE nodes ADD COLUMN relay_version text`,
`ALTER TABLE nodes ADD COLUMN relay_frp_version text`,
`ALTER TABLE nodes ADD COLUMN relay_frps_connections integer`,
`ALTER TABLE nodes ADD COLUMN relay_frps_proxy_count integer`,
} {
if err := db.Exec(stmt).Error; err != nil {
t.Fatalf("prepare legacy node column with %q: %v", stmt, err)
}
}
now := time.Now()
if err := db.Exec(`
INSERT INTO nodes (
node_id, name, ip, access_token, version, ext_version,
agent_token, agent_version, nginx_version,
status, last_seen_at, created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, "node-v16", "Node v16", "127.0.0.1", "", "", "", "legacy-token", "v2.0.0", "openresty/1.25.3", "offline", now, now, now).Error; err != nil {
t.Fatalf("seed node with legacy columns: %v", err)
}
if err := saveDatabaseSchemaVersion(db, 15); err != nil {
t.Fatalf("save schema version: %v", err)
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
}
var node Node
if err := db.Where("node_id = ?", "node-v16").First(&node).Error; err != nil {
t.Fatalf("query migrated node: %v", err)
}
if node.AccessToken != "legacy-token" {
t.Fatalf("unexpected access_token: got %q", node.AccessToken)
}
if node.Version != "v2.0.0" {
t.Fatalf("unexpected version: got %q", node.Version)
}
if node.ExtVersion != "openresty/1.25.3" {
t.Fatalf("unexpected ext_version: got %q", node.ExtVersion)
}
for _, column := range []string{
"agent_token",
"agent_version",
"nginx_version",
"relay_version",
"relay_frp_version",
"relay_frps_connections",
"relay_frps_proxy_count",
} {
exists, err := databaseColumnExists(db, "nodes", column)
if err != nil {
t.Fatalf("inspect legacy nodes.%s: %v", column, err)
}
if exists {
t.Fatalf("expected migration to drop legacy nodes.%s column", column)
}
}
version, exists, err := loadDatabaseSchemaVersion(db)
if err != nil {
t.Fatalf("loadDatabaseSchemaVersion: %v", err)
}
if !exists {
t.Fatal("expected schema version record to exist")
}
if version != expectedCurrentDatabaseVersion() {
t.Fatalf("unexpected schema version: got %d want %d", version, expectedCurrentDatabaseVersion())
}
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
t.Fatal("expected v16 upgrade path to apply goose migration nodes.capabilities_json")
}
}
func TestEnsureDatabaseSchemaUpToDateV16DropsLegacyNodeColumnsWhenAlreadyCurrent(t *testing.T) {
db := openBareTestSQLiteDB(t, "node-v16-current-legacy-columns.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := applyCurrentSchema(db, "sqlite"); err != nil {
t.Fatalf("apply current schema: %v", err)
}
for _, stmt := range []string{
`ALTER TABLE nodes ADD COLUMN agent_token text`,
`ALTER TABLE nodes ADD COLUMN agent_version text`,
`ALTER TABLE nodes ADD COLUMN nginx_version text`,
} {
if err := db.Exec(stmt).Error; err != nil {
t.Fatalf("prepare legacy node column with %q: %v", stmt, err)
}
}
if err := saveDatabaseSchemaVersion(db, currentDatabaseSchemaVersion); err != nil {
t.Fatalf("save schema version: %v", err)
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
}
for _, column := range []string{"agent_token", "agent_version", "nginx_version"} {
exists, err := databaseColumnExists(db, "nodes", column)
if err != nil {
t.Fatalf("inspect legacy nodes.%s: %v", column, err)
}
if exists {
t.Fatalf("expected current-schema cleanup to drop legacy nodes.%s column", column)
}
}
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
t.Fatal("expected current-schema legacy version table to be removed after goose bridge")
}
if !db.Migrator().HasTable("goose_db_version") {
t.Fatal("expected current-schema goose_db_version table to exist")
}
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
t.Fatal("expected current-schema repair to preserve goose column nodes.capabilities_json")
}
}
func TestEnsureDatabaseSchemaUpToDateKeepsGooseOnlyDatabaseOnReentry(t *testing.T) {
db := openBareTestSQLiteDB(t, "goose-only-reentry.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("first ensureDatabaseSchemaUpToDate: %v", err)
}
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
t.Fatal("expected first initialization to avoid legacy table")
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("second ensureDatabaseSchemaUpToDate: %v", err)
}
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
t.Fatal("expected goose-only database to remain free of legacy version table")
}
version, exists, err := loadGooseDatabaseVersion(db)
if err != nil {
t.Fatalf("loadGooseDatabaseVersion: %v", err)
}
if !exists {
t.Fatal("expected goose-only database to keep goose version record")
}
if version != expectedCurrentDatabaseVersion() {
t.Fatalf("unexpected goose version: got %d want %d", version, expectedCurrentDatabaseVersion())
}
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
t.Fatal("expected goose-only database to keep nodes.capabilities_json")
}
}
func TestAllRegisteredMigrationsHaveValidationDefined(t *testing.T) {
ctx := databaseSchemaMigrationContext{}
for _, migration := range databaseSchemaMigrations() {
err := ctx.ValidateDatabaseSchemaVersion(nil, "sqlite", migration.toVersion)
if err != nil && strings.Contains(err.Error(), "is not defined") {
t.Fatalf("Validation is not defined in migrations.go for registered migration version v%d: %v", migration.toVersion, err)
}
}
}
func TestAllGORMModelsAreRegistered(t *testing.T) {
// 1. Gather all registered model names
registeredNames := make(map[string]bool)
for _, item := range registeredModels() {
name := reflect.TypeOf(item).Elem().Name()
registeredNames[name] = true
}
for _, item := range schemaMetadataModels() {
name := reflect.TypeOf(item).Elem().Name()
registeredNames[name] = true
}
// 2. Parse all .go files in model/ package
fset := token.NewFileSet()
pkgs, err := parser.ParseDir(fset, ".", func(info os.FileInfo) bool {
// Only parse .go files, exclude _test.go files and subdirectories
return !info.IsDir() && strings.HasSuffix(info.Name(), ".go") && !strings.HasSuffix(info.Name(), "_test.go")
}, 0)
if err != nil {
t.Fatalf("failed to parse directory: %v", err)
}
for _, pkg := range pkgs {
for _, file := range pkg.Files {
for _, decl := range file.Decls {
genDecl, ok := decl.(*ast.GenDecl)
if !ok || genDecl.Tok != token.TYPE {
continue
}
for _, spec := range genDecl.Specs {
typeSpec, ok := spec.(*ast.TypeSpec)
if !ok {
continue
}
structType, ok := typeSpec.Type.(*ast.StructType)
if !ok {
continue
}
// Verify if this struct has any field with a `gorm:"..."` tag
isGORMModel := false
for _, field := range structType.Fields.List {
if field.Tag != nil && strings.Contains(field.Tag.Value, "gorm:") {
isGORMModel = true
break
}
}
if isGORMModel {
structName := typeSpec.Name.Name
if !registeredNames[structName] {
t.Errorf("Model struct %q is defined with GORM tags but is NOT registered in registeredModels() or schemaMetadataModels() in model/main.go!", structName)
}
}
}
}
}
}
}
func TestEnsureDatabaseSchemaUpToDateDropsPagesDeploymentUnusedFields(t *testing.T) {
db := openBareTestSQLiteDB(t, "drop-pages-deployment-unused-fields.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("first ensureDatabaseSchemaUpToDate: %v", err)
}
// Verify columns do not exist
if db.Migrator().HasColumn("pages_deployments", "root_dir") {
t.Fatal("expected root_dir column to be absent initially")
}
if db.Migrator().HasColumn("pages_deployments", "entry_file") {
t.Fatal("expected entry_file column to be absent initially")
}
// Manually add columns to simulate old state
if err := db.Exec("ALTER TABLE pages_deployments ADD COLUMN root_dir TEXT").Error; err != nil {
t.Fatalf("failed to add root_dir column: %v", err)
}
if err := db.Exec("ALTER TABLE pages_deployments ADD COLUMN entry_file TEXT").Error; err != nil {
t.Fatalf("failed to add entry_file column: %v", err)
}
// Verify columns were added
if !db.Migrator().HasColumn("pages_deployments", "root_dir") {
t.Fatal("expected root_dir column to be present after manual add")
}
if !db.Migrator().HasColumn("pages_deployments", "entry_file") {
t.Fatal("expected entry_file column to be present after manual add")
}
// Remove the migration record from goose_db_version table
const versionToRerun = 202606040004
if err := db.Exec("DELETE FROM goose_db_version WHERE version_id = ?", versionToRerun).Error; err != nil {
t.Fatalf("failed to delete migration record: %v", err)
}
// Run migration again
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("second ensureDatabaseSchemaUpToDate: %v", err)
}
// Verify columns were dropped successfully
if db.Migrator().HasColumn("pages_deployments", "root_dir") {
t.Fatal("expected root_dir column to be dropped after migration rerun")
}
if db.Migrator().HasColumn("pages_deployments", "entry_file") {
t.Fatal("expected entry_file column to be dropped after migration rerun")
}
}
@@ -0,0 +1,41 @@
package model
import "time"
type ManagedDomain struct {
ID uint `json:"id" gorm:"primaryKey"`
Domain string `json:"domain" gorm:"uniqueIndex;size:255;not null"`
CertID *uint `json:"cert_id"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
Remark string `json:"remark" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func ListManagedDomains() (domains []*ManagedDomain, err error) {
err = DB.Order("id desc").Find(&domains).Error
return domains, err
}
func ListEnabledManagedDomainsWithCertificate() (domains []*ManagedDomain, err error) {
err = DB.Where("enabled = ? AND cert_id IS NOT NULL", true).Order("id desc").Find(&domains).Error
return domains, err
}
func GetManagedDomainByID(id uint) (*ManagedDomain, error) {
domain := &ManagedDomain{}
err := DB.First(domain, id).Error
return domain, err
}
func (domain *ManagedDomain) Insert() error {
return DB.Create(domain).Error
}
func (domain *ManagedDomain) Update() error {
return DB.Save(domain).Error
}
func (domain *ManagedDomain) Delete() error {
return DB.Delete(domain).Error
}
@@ -0,0 +1,4 @@
package migrate
// Versions 1 through 7 are treated as the historical baseline. There are no
// supported deployments below v8, so new upgrades start from this base version.
@@ -0,0 +1,54 @@
package migrate
import (
"sort"
"gorm.io/gorm"
)
const BaseDatabaseSchemaVersion = 7
type Context interface {
ApplyCurrentSchema(db *gorm.DB, backend string) error
ApplyCurrentSchemaExcept(db *gorm.DB, backend string, excludedTables ...string) error
BackfillOriginsFromProxyRoutes(db *gorm.DB) error
BackfillProxyRouteSiteFields(db *gorm.DB) error
EnsureProxyRouteSiteNameUniqueIndex(db *gorm.DB) error
BackfillProxyRouteCertificateFields(db *gorm.DB) error
BackfillProxyRouteDomainCertificateFields(db *gorm.DB) error
EnsureDefaultGitHubAuthSource(db *gorm.DB) error
EnsureDefaultWAFRuleGroup(db *gorm.DB) error
DropLegacyNodeColumns(db *gorm.DB, backend string) error
ValidateDatabaseSchemaVersion(db *gorm.DB, backend string, version int) error
}
type Migration struct {
FromVersion int
ToVersion int
Migrate func(ctx Context, db *gorm.DB, backend string) error
Validate func(ctx Context, db *gorm.DB, backend string) error
}
var registeredMigrations []Migration
func Register(migration Migration) {
registeredMigrations = append(registeredMigrations, migration)
}
func Migrations() []Migration {
migrations := append([]Migration{}, registeredMigrations...)
sort.Slice(migrations, func(i int, j int) bool {
return migrations[i].FromVersion < migrations[j].FromVersion
})
return migrations
}
func CurrentVersion() int {
version := BaseDatabaseSchemaVersion
for _, migration := range registeredMigrations {
if migration.ToVersion > version {
version = migration.ToVersion
}
}
return version
}
@@ -0,0 +1,23 @@
package migrate
import "testing"
func TestMigrationsAreContinuousFromBaseVersion(t *testing.T) {
migrations := Migrations()
if len(migrations) == 0 {
t.Fatal("expected at least one registered migration")
}
expectedFrom := BaseDatabaseSchemaVersion
for _, migration := range migrations {
if migration.FromVersion != expectedFrom {
t.Fatalf("expected migration from v%d, got v%d -> v%d", expectedFrom, migration.FromVersion, migration.ToVersion)
}
if migration.ToVersion != migration.FromVersion+1 {
t.Fatalf("expected one-step migration, got v%d -> v%d", migration.FromVersion, migration.ToVersion)
}
expectedFrom = migration.ToVersion
}
if CurrentVersion() != expectedFrom {
t.Fatalf("unexpected current version: got %d want %d", CurrentVersion(), expectedFrom)
}
}
@@ -0,0 +1,26 @@
// v10 升级内容:新增可配置认证源与第三方账号绑定,并迁移旧 GitHub 登录配置。
// 背景说明:登录体系从固定 GitHub OAuth 字段演进为通用认证源模型,需要创建 auth_sources、external_accounts,并把旧用户 GitHub 绑定迁移到新表。
package migrate
import "gorm.io/gorm"
func init() {
Register(V10())
}
func V10() Migration {
return Migration{
FromVersion: 9,
ToVersion: 10,
Migrate: migrateV10,
Validate: validateV10,
}
}
func migrateV10(ctx Context, db *gorm.DB, backend string) error {
return ctx.EnsureDefaultGitHubAuthSource(db)
}
func validateV10(ctx Context, db *gorm.DB, backend string) error {
return ctx.ValidateDatabaseSchemaVersion(db, backend, 10)
}
@@ -0,0 +1,26 @@
// v11 升级内容:新增 ACME 账户、DNS 账户,并扩展证书 provider 字段。
// 背景说明:证书申请能力从单一手工导入扩展到自动签发,需要持久化 ACME/DNS 凭据,并标记证书来源。
package migrate
import "gorm.io/gorm"
func init() {
Register(V11())
}
func V11() Migration {
return Migration{
FromVersion: 10,
ToVersion: 11,
Migrate: migrateV11,
Validate: validateV11,
}
}
func migrateV11(ctx Context, db *gorm.DB, backend string) error {
return nil
}
func validateV11(ctx Context, db *gorm.DB, backend string) error {
return ctx.ValidateDatabaseSchemaVersion(db, backend, 11)
}
@@ -0,0 +1,26 @@
// v12 升级内容:为 proxy_routes 增加 Basic Auth 相关字段。
// 背景说明:站点级访问控制需要支持基础认证,因此在代理路由配置中持久化 Basic Auth 开关与凭据配置。
package migrate
import "gorm.io/gorm"
func init() {
Register(V12())
}
func V12() Migration {
return Migration{
FromVersion: 11,
ToVersion: 12,
Migrate: migrateV12,
Validate: validateV12,
}
}
func migrateV12(ctx Context, db *gorm.DB, backend string) error {
return nil
}
func validateV12(ctx Context, db *gorm.DB, backend string) error {
return ctx.ValidateDatabaseSchemaVersion(db, backend, 12)
}
@@ -0,0 +1,26 @@
// v13 升级内容:新增 WAF 规则组与站点绑定表,并创建默认全局规则组。
// 背景说明:WAF 配置从零散站点字段演进为可复用规则组,需要全局规则组作为默认入口,并支持站点与规则组绑定。
package migrate
import "gorm.io/gorm"
func init() {
Register(V13())
}
func V13() Migration {
return Migration{
FromVersion: 12,
ToVersion: 13,
Migrate: migrateV13,
Validate: validateV13,
}
}
func migrateV13(ctx Context, db *gorm.DB, backend string) error {
return ctx.EnsureDefaultWAFRuleGroup(db)
}
func validateV13(ctx Context, db *gorm.DB, backend string) error {
return ctx.ValidateDatabaseSchemaVersion(db, backend, 13)
}
@@ -0,0 +1,26 @@
// v14 升级内容:为 WAF 规则组增加 PoW 策略字段。
// 背景说明:PoW 能力从站点路由侧沉淀到 WAF 规则组中,便于统一按规则组管理人机挑战策略。
package migrate
import "gorm.io/gorm"
func init() {
Register(V14())
}
func V14() Migration {
return Migration{
FromVersion: 13,
ToVersion: 14,
Migrate: migrateV14,
Validate: validateV14,
}
}
func migrateV14(ctx Context, db *gorm.DB, backend string) error {
return ctx.EnsureDefaultWAFRuleGroup(db)
}
func validateV14(ctx Context, db *gorm.DB, backend string) error {
return ctx.ValidateDatabaseSchemaVersion(db, backend, 14)
}
@@ -0,0 +1,52 @@
// v15 升级内容:为 nodes 增加 ip_manual_override 字段。
// 背景说明:管理端手动指定节点 IP 后,Agent 心跳不应继续覆盖该值,因此需要在节点表中记录 IP 是否由管理端锁定。
package migrate
import (
"fmt"
"gorm.io/gorm"
)
type nodeV15 struct {
IPManualOverride bool `gorm:"column:ip_manual_override;not null;default:false"`
}
func init() {
Register(V15())
}
func V15() Migration {
return Migration{
FromVersion: 14,
ToVersion: 15,
Migrate: migrateV15,
Validate: validateV15,
}
}
func (nodeV15) TableName() string {
return "nodes"
}
func migrateV15(ctx Context, db *gorm.DB, backend string) error {
if db == nil {
return fmt.Errorf("database handle is nil")
}
if !db.Migrator().HasColumn(&nodeV15{}, "ip_manual_override") {
if err := db.Migrator().AddColumn(&nodeV15{}, "IPManualOverride"); err != nil {
return fmt.Errorf("add nodes.ip_manual_override: %w", err)
}
}
return nil
}
func validateV15(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ValidateDatabaseSchemaVersion(db, backend, 14); err != nil {
return err
}
if db == nil || !db.Migrator().HasColumn(&nodeV15{}, "ip_manual_override") {
return fmt.Errorf("column nodes.ip_manual_override is missing")
}
return nil
}
@@ -0,0 +1,185 @@
// v16 is the first database migration after the V15 formal release baseline.
// It folds the previously drafted v16-v21 schema work into a single official
// upgrade: tunnel-relay fields, WAF IP groups, current node identity/version
// columns, and split node observation tables. The migration also backfills
// legacy node columns and removes obsolete pre-release tunnel metadata when
// present, so V15 deployments can upgrade directly to the new formal schema.
package migrate
import (
"fmt"
"log/slog"
"gorm.io/gorm"
)
func (nodeV16) TableName() string {
return "nodes"
}
func (tunnelV16) TableName() string {
return "tunnels"
}
func (proxyRouteV16) TableName() string {
return "proxy_routes"
}
type nodeV16 struct{}
type tunnelV16 struct{}
type proxyRouteV16 struct{}
type wafIPGroupV16 struct{}
type wafRuleGroupV16 struct{}
func (wafIPGroupV16) TableName() string {
return "waf_ip_groups"
}
func (wafRuleGroupV16) TableName() string {
return "waf_rule_groups"
}
func init() {
Register(V16())
}
func V16() Migration {
return Migration{
FromVersion: 15,
ToVersion: 16,
Migrate: migrateV16,
Validate: validateV16,
}
}
func migrateV16(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
migrator := db.Migrator()
if migrator.HasColumn(&nodeV16{}, "agent_token") {
if err := db.Exec(`UPDATE nodes SET access_token = agent_token WHERE access_token IS NULL OR access_token = ''`).Error; err != nil {
return fmt.Errorf("backfill nodes.access_token from agent_token: %w", err)
}
}
if migrator.HasColumn(&nodeV16{}, "agent_version") {
if err := db.Exec(`UPDATE nodes SET version = agent_version WHERE version = '' OR version IS NULL`).Error; err != nil {
return fmt.Errorf("backfill nodes.version from agent_version: %w", err)
}
}
if migrator.HasColumn(&nodeV16{}, "nginx_version") {
if err := db.Exec(`UPDATE nodes SET ext_version = nginx_version WHERE ext_version IS NULL OR ext_version = ''`).Error; err != nil {
return fmt.Errorf("backfill nodes.ext_version from nginx_version: %w", err)
}
}
if err := ctx.DropLegacyNodeColumns(db, backend); err != nil {
return err
}
if err := db.Exec("UPDATE nodes SET node_type = 'edge_node' WHERE node_type = '' OR node_type IS NULL").Error; err != nil {
return fmt.Errorf("backfill nodes.node_type: %w", err)
}
if err := db.Exec("UPDATE proxy_routes SET upstream_type = 'direct' WHERE upstream_type = '' OR upstream_type IS NULL").Error; err != nil {
return fmt.Errorf("backfill proxy_routes.upstream_type: %w", err)
}
if migrator.HasColumn(&proxyRouteV16{}, "tunnel_id") {
if err := db.Model(&proxyRouteV16{}).Where("upstream_type = ?", "tunnel").Update("upstream_type", "direct").Error; err != nil {
return fmt.Errorf("reset pre-release tunnel proxy routes: %w", err)
}
// Drop the legacy index idx_proxy_routes_tunnel_id if it exists, to avoid errors on dropping the tunnel_id column (especially on SQLite).
if migrator.HasIndex(&proxyRouteV16{}, "idx_proxy_routes_tunnel_id") {
if err := migrator.DropIndex(&proxyRouteV16{}, "idx_proxy_routes_tunnel_id"); err != nil {
return fmt.Errorf("drop index idx_proxy_routes_tunnel_id failed: %w", err)
}
}
if err := migrator.DropColumn(&proxyRouteV16{}, "tunnel_id"); err != nil {
return fmt.Errorf("drop pre-release proxy_routes.tunnel_id: %w", err)
}
}
if migrator.HasTable(&tunnelV16{}) {
if err := migrator.DropTable(&tunnelV16{}); err != nil {
return fmt.Errorf("drop pre-release tunnels table: %w", err)
}
slog.Info("dropped pre-release tunnels table during v16 migration")
}
return nil
}
func validateV16(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ValidateDatabaseSchemaVersion(db, backend, 15); err != nil {
return err
}
if db == nil {
return fmt.Errorf("database handle is nil")
}
migrator := db.Migrator()
for _, column := range []string{
"access_token",
"version",
"ext_version",
"node_type",
"relay_bind_port",
"relay_vhost_http_port",
"relay_auth_token",
"relay_agent_access_addr",
"relay_client_access_addr",
"relay_client_proxy_url",
"relay_status",
} {
if !migrator.HasColumn(&nodeV16{}, column) {
return fmt.Errorf("column nodes.%s is missing", column)
}
}
for _, column := range []string{
"upstream_type",
"tunnel_node_id",
"tunnel_target_addr",
"tunnel_target_protocol",
} {
if !migrator.HasColumn(&proxyRouteV16{}, column) {
return fmt.Errorf("column proxy_routes.%s is missing", column)
}
}
if migrator.HasColumn(&proxyRouteV16{}, "tunnel_id") {
return fmt.Errorf("column proxy_routes.tunnel_id should not exist in v16")
}
if migrator.HasTable(&tunnelV16{}) {
return fmt.Errorf("table tunnels should not exist in v16")
}
for _, column := range []string{
"agent_token",
"agent_version",
"nginx_version",
"relay_version",
"relay_frp_version",
"relay_frps_connections",
"relay_frps_proxy_count",
} {
if migrator.HasColumn(&nodeV16{}, column) {
return fmt.Errorf("column nodes.%s should not exist in v16", column)
}
}
if !migrator.HasTable(&wafIPGroupV16{}) {
return fmt.Errorf("table waf_ip_groups is missing")
}
for _, column := range []string{
"ip_whitelist_groups",
"ip_blacklist_groups",
} {
if !migrator.HasColumn(&wafRuleGroupV16{}, column) {
return fmt.Errorf("column waf_rule_groups.%s is missing", column)
}
}
if !migrator.HasColumn(&wafIPGroupV16{}, "ext_ips") {
return fmt.Errorf("column waf_ip_groups.ext_ips is missing")
}
return nil
}
@@ -0,0 +1,58 @@
package migrate
import (
"fmt"
"gorm.io/gorm"
)
type nodeV17 struct{}
func (nodeV17) TableName() string {
return "nodes"
}
func init() {
Register(V17())
}
func V17() Migration {
return Migration{
FromVersion: 16,
ToVersion: 17,
Migrate: migrateV17,
Validate: validateV17,
}
}
func migrateV17(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
return nil
}
func validateV17(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ValidateDatabaseSchemaVersion(db, backend, 16); err != nil {
return err
}
if db == nil {
return fmt.Errorf("database handle is nil")
}
migrator := db.Migrator()
if !migrator.HasColumn(&nodeV17{}, "relay_web_server_enabled") {
return fmt.Errorf("column nodes.relay_web_server_enabled is missing")
}
// Validate columns on a sharded partition table
for _, shard := range []string{"node_observation_frps_00"} {
for _, column := range []string{"frps_client_count", "frps_proxies"} {
if !migrator.HasColumn(shard, column) {
return fmt.Errorf("column %s.%s is missing", shard, column)
}
}
}
return nil
}
@@ -0,0 +1,41 @@
// v8 升级内容:为 proxy_routes 增加域名级证书绑定字段 domain_cert_ids,并回填已有站点的证书映射。
// 背景说明:v1-v7 已作为历史初始基线合并;v8 是当前保留逐版本升级链的起点,用于把早期站点级证书列表扩展为每个域名可独立绑定证书。
package migrate
import "gorm.io/gorm"
func init() {
Register(V8())
}
func V8() Migration {
return Migration{
FromVersion: 7,
ToVersion: 8,
Migrate: migrateV8,
Validate: validateV8,
}
}
func migrateV8(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
if err := ctx.BackfillOriginsFromProxyRoutes(db); err != nil {
return err
}
if err := ctx.BackfillProxyRouteSiteFields(db); err != nil {
return err
}
if err := ctx.EnsureProxyRouteSiteNameUniqueIndex(db); err != nil {
return err
}
if err := ctx.BackfillProxyRouteCertificateFields(db); err != nil {
return err
}
return ctx.BackfillProxyRouteDomainCertificateFields(db)
}
func validateV8(ctx Context, db *gorm.DB, backend string) error {
return ctx.ValidateDatabaseSchemaVersion(db, backend, 8)
}
@@ -0,0 +1,29 @@
// v9 升级内容:为 proxy_routes 增加 PoW 防护配置字段。
// 背景说明:反向代理站点需要支持 Proof-of-Work 抗机器人能力,因此在路由配置中持久化 PoW 开关与策略,并沿用 v8 的证书与站点字段回填。
package migrate
import "gorm.io/gorm"
func init() {
Register(V9())
}
func V9() Migration {
return Migration{
FromVersion: 8,
ToVersion: 9,
Migrate: migrateV9,
Validate: validateV9,
}
}
func migrateV9(ctx Context, db *gorm.DB, backend string) error {
if err := migrateV8(ctx, db, backend); err != nil {
return err
}
return nil
}
func validateV9(ctx Context, db *gorm.DB, backend string) error {
return ctx.ValidateDatabaseSchemaVersion(db, backend, 9)
}
File diff suppressed because it is too large Load Diff
+91
View File
@@ -0,0 +1,91 @@
package model
import "time"
type Node struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"uniqueIndex;size:64;not null"`
Name string `json:"name" gorm:"size:128;not null"`
IP string `json:"ip" gorm:"size:64;not null"`
IPManualOverride bool `json:"ip_manual_override" gorm:"not null;default:false"`
GeoName string `json:"geo_name" gorm:"size:128"`
GeoLatitude *float64 `json:"geo_latitude"`
GeoLongitude *float64 `json:"geo_longitude"`
GeoManualOverride bool `json:"geo_manual_override" gorm:"not null;default:false"`
AccessToken string `json:"-" gorm:"column:access_token;size:128;index"`
AutoUpdateEnabled bool `json:"auto_update_enabled" gorm:"not null;default:false"`
UpdateRequested bool `json:"update_requested" gorm:"not null;default:false"`
UpdateChannel string `json:"update_channel" gorm:"size:16;not null;default:'stable'"`
UpdateTag string `json:"update_tag" gorm:"size:64"`
RestartOpenrestyRequested bool `json:"restart_openresty_requested" gorm:"not null;default:false"`
Version string `json:"version" gorm:"size:64;not null;default:''"`
ExtVersion string `json:"ext_version" gorm:"size:64"`
OpenrestyStatus string `json:"openresty_status" gorm:"size:16;not null;default:'unknown'"`
OpenrestyMessage string `json:"openresty_message" gorm:"type:text"`
Status string `json:"status" gorm:"size:16;not null;default:'offline'"`
CurrentVersion string `json:"current_version" gorm:"size:32"`
LastSeenAt time.Time `json:"last_seen_at"`
LastError string `json:"last_error" gorm:"type:text"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
// Node type: edge_node (default) | tunnel_relay | tunnel_client
NodeType string `json:"node_type" gorm:"size:32;not null;default:'edge_node'"`
// TunnelRelay specific fields
RelayBindPort int `json:"relay_bind_port" gorm:"not null;default:0"`
RelayVhostHTTPPort int `json:"relay_vhost_http_port" gorm:"not null;default:0"`
RelayAuthToken string `json:"-" gorm:"size:128"`
RelayAgentAccessAddr string `json:"relay_agent_access_addr" gorm:"size:255"`
RelayClientAccessAddr string `json:"relay_client_access_addr" gorm:"size:255"`
RelayClientProxyURL string `json:"relay_client_proxy_url" gorm:"size:512"`
CapabilitiesJSON string `json:"capabilities_json" gorm:"type:text;not null;default:'[]'"`
RelayStatus string `json:"relay_status" gorm:"size:16;not null;default:'unknown'"`
RelayWebServerEnabled bool `json:"relay_web_server_enabled" gorm:"not null;default:false"`
}
func ListNodes() (nodes []*Node, err error) {
err = DB.Order("id desc").Find(&nodes).Error
return nodes, err
}
func ListNodesByNodeIDs(nodeIDs []string) (nodes []*Node, err error) {
if len(nodeIDs) == 0 {
return []*Node{}, nil
}
err = DB.Where("node_id IN ?", nodeIDs).Find(&nodes).Error
return nodes, err
}
func GetNodeByNodeID(nodeID string) (*Node, error) {
node := &Node{}
err := DB.Where("node_id = ?", nodeID).First(node).Error
return node, err
}
func GetNodeByID(id uint) (*Node, error) {
node := &Node{}
err := DB.First(node, id).Error
return node, err
}
func GetNodeByAccessToken(token string) (*Node, error) {
node := &Node{}
err := DB.Where("access_token = ?", token).First(node).Error
return node, err
}
func (node *Node) Insert() error {
return DB.Create(node).Error
}
func (node *Node) Update() error {
return DB.Save(node).Error
}
func (node *Node) Delete() error {
return DB.Delete(node).Error
}
func ListNodesByType(nodeType string) (nodes []*Node, err error) {
err = DB.Where("node_type = ?", nodeType).Order("id desc").Find(&nodes).Error
return nodes, err
}
@@ -0,0 +1,749 @@
package model
import (
"sort"
"strings"
"time"
"gorm.io/gorm"
)
type NodeAccessLog struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index:,composite:node_logged_at,priority:1;size:64;not null"`
LoggedAt time.Time `json:"logged_at" gorm:"index;index:,composite:node_logged_at,priority:2"`
RemoteAddr string `json:"remote_addr" gorm:"index;size:128"`
Region string `json:"region" gorm:"size:128"`
Host string `json:"host" gorm:"index;size:255"`
Path string `json:"path" gorm:"size:2048"`
StatusCode int `json:"status_code" gorm:"index"`
CreatedAt time.Time `json:"created_at"`
}
type NodeAccessLogRegionCount struct {
Region string `json:"region"`
Count int64 `json:"count"`
}
type NodeAccessLogQuery struct {
NodeID string
RemoteAddr string
Host string
Path string
Since time.Time
Until time.Time
Page int
PageSize int
SortBy string
SortOrder string
}
type NodeAccessLogBucketQuery struct {
NodeID string
RemoteAddr string
Host string
Path string
Since time.Time
Page int
PageSize int
SortBy string
SortOrder string
FoldMinutes int
}
type NodeAccessLogBucketRow struct {
BucketEpoch int64 `json:"bucket_epoch"`
RequestCount int64 `json:"request_count"`
UniqueIPCount int64 `json:"unique_ip_count"`
UniqueHostCount int64 `json:"unique_host_count"`
SuccessCount int64 `json:"success_count"`
ClientErrorCount int64 `json:"client_error_count"`
ServerErrorCount int64 `json:"server_error_count"`
}
type NodeAccessLogBucketIPQuery struct {
NodeID string
RemoteAddr string
Host string
Path string
BucketStartedAt time.Time
FoldMinutes int
Page int
PageSize int
SortBy string
SortOrder string
}
type NodeAccessLogBucketIPRow struct {
RemoteAddr string `json:"remote_addr"`
RequestCount int64 `json:"request_count"`
SuccessCount int64 `json:"success_count"`
ClientErrorCount int64 `json:"client_error_count"`
ServerErrorCount int64 `json:"server_error_count"`
LastSeenEpoch int64 `json:"last_seen_epoch"`
}
type NodeAccessLogIPSummaryQuery struct {
NodeID string
RemoteAddr string
Host string
Since time.Time
Page int
PageSize int
SortBy string
SortOrder string
}
type NodeAccessLogIPSummaryRow struct {
RemoteAddr string `json:"remote_addr"`
TotalRequests int64 `json:"total_requests"`
RecentRequests int64 `json:"recent_requests"`
LastSeenEpoch int64 `json:"last_seen_epoch"`
}
type NodeAccessLogIPTrendQuery struct {
NodeID string
RemoteAddr string
Host string
Since time.Time
BucketMinutes int
}
type NodeAccessLogTrendPointRow struct {
BucketEpoch int64 `json:"bucket_epoch"`
RequestCount int64 `json:"request_count"`
}
func (log *NodeAccessLog) BeforeCreate(*gorm.DB) error {
return assignObservabilityID(&log.ID)
}
func ListNodeAccessLogs(query NodeAccessLogQuery) (logs []*NodeAccessLog, err error) {
all, err := listNodeAccessLogsAcrossShards(query)
if err != nil {
return nil, err
}
start, end := paginateBounds(len(all), query.Page, query.PageSize)
if start >= len(all) {
return []*NodeAccessLog{}, nil
}
return all[start:end], nil
}
func ListNodeAccessLogsForWAFIPGroup(query NodeAccessLogQuery) ([]*NodeAccessLog, error) {
return listNodeAccessLogsAcrossShards(query)
}
func CountNodeAccessLogs(query NodeAccessLogQuery) (totalRecords int64, totalIPs int64, err error) {
all, err := listNodeAccessLogsAcrossShards(query)
if err != nil {
return 0, 0, err
}
ips := make(map[string]struct{}, len(all))
for _, item := range all {
if item == nil {
continue
}
trimmed := strings.TrimSpace(item.RemoteAddr)
if trimmed != "" {
ips[trimmed] = struct{}{}
}
}
return int64(len(all)), int64(len(ips)), nil
}
func ListNodeAccessLogRegionCounts(nodeID string, since time.Time, limit int) (items []*NodeAccessLogRegionCount, err error) {
logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{
NodeID: nodeID,
Since: since,
})
if err != nil {
return nil, err
}
counts := make(map[string]int64)
for _, item := range logs {
if item == nil {
continue
}
region := strings.TrimSpace(item.Region)
if region == "" {
continue
}
counts[region]++
}
items = make([]*NodeAccessLogRegionCount, 0, len(counts))
for region, count := range counts {
items = append(items, &NodeAccessLogRegionCount{
Region: region,
Count: count,
})
}
sort.Slice(items, func(i int, j int) bool {
if items[i].Count == items[j].Count {
return items[i].Region < items[j].Region
}
return items[i].Count > items[j].Count
})
if limit > 0 && len(items) > limit {
items = items[:limit]
}
return items, nil
}
func ListNodeAccessLogBuckets(query NodeAccessLogBucketQuery) (items []*NodeAccessLogBucketRow, err error) {
rows, err := buildNodeAccessLogBucketRows(query)
if err != nil {
return nil, err
}
start, end := paginateBounds(len(rows), query.Page, query.PageSize)
if start >= len(rows) {
return []*NodeAccessLogBucketRow{}, nil
}
return rows[start:end], nil
}
func CountNodeAccessLogBuckets(query NodeAccessLogBucketQuery) (total int64, err error) {
rows, err := buildNodeAccessLogBucketRows(query)
if err != nil {
return 0, err
}
return int64(len(rows)), nil
}
func ListNodeAccessLogBucketIPs(query NodeAccessLogBucketIPQuery) (items []*NodeAccessLogBucketIPRow, err error) {
rows, err := buildNodeAccessLogBucketIPRows(query)
if err != nil {
return nil, err
}
start, end := paginateBounds(len(rows), query.Page, query.PageSize)
if start >= len(rows) {
return []*NodeAccessLogBucketIPRow{}, nil
}
return rows[start:end], nil
}
func CountNodeAccessLogBucketIPs(query NodeAccessLogBucketIPQuery) (total int64, err error) {
rows, err := buildNodeAccessLogBucketIPRows(query)
if err != nil {
return 0, err
}
return int64(len(rows)), nil
}
func ListNodeAccessLogIPSummaries(query NodeAccessLogIPSummaryQuery, recentSince time.Time) (items []*NodeAccessLogIPSummaryRow, err error) {
rows, err := buildNodeAccessLogIPSummaryRows(query, recentSince)
if err != nil {
return nil, err
}
start, end := paginateBounds(len(rows), query.Page, query.PageSize)
if start >= len(rows) {
return []*NodeAccessLogIPSummaryRow{}, nil
}
return rows[start:end], nil
}
func CountNodeAccessLogIPSummaries(query NodeAccessLogIPSummaryQuery) (total int64, err error) {
rows, err := buildNodeAccessLogIPSummaryRows(query, time.Time{})
if err != nil {
return 0, err
}
return int64(len(rows)), nil
}
func ListNodeAccessLogIPTrend(query NodeAccessLogIPTrendQuery) (items []*NodeAccessLogTrendPointRow, err error) {
logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: query.RemoteAddr,
Host: query.Host,
Since: query.Since,
})
if err != nil {
return nil, err
}
remoteAddr := strings.TrimSpace(query.RemoteAddr)
if remoteAddr == "" {
return []*NodeAccessLogTrendPointRow{}, nil
}
buckets := make(map[int64]int64)
for _, item := range logs {
if item == nil || strings.TrimSpace(item.RemoteAddr) != remoteAddr {
continue
}
bucketEpoch := bucketEpochForTime(item.LoggedAt, query.BucketMinutes)
buckets[bucketEpoch]++
}
items = make([]*NodeAccessLogTrendPointRow, 0, len(buckets))
for bucketEpoch, requestCount := range buckets {
items = append(items, &NodeAccessLogTrendPointRow{
BucketEpoch: bucketEpoch,
RequestCount: requestCount,
})
}
sort.Slice(items, func(i int, j int) bool {
return items[i].BucketEpoch < items[j].BucketEpoch
})
return items, nil
}
func DeleteNodeAccessLogsBefore(before time.Time) (deleted int64, err error) {
return deleteAcrossShards(DB, "node_access_logs", &NodeAccessLog{}, func(tx *gorm.DB) *gorm.DB {
return tx.Where("logged_at < ?", before)
})
}
func DeleteAllNodeAccessLogs(db *gorm.DB) (deleted int64, err error) {
return deleteAcrossShards(db, "node_access_logs", &NodeAccessLog{}, nil)
}
func NodeAccessLogExists(db *gorm.DB, record *NodeAccessLog) (bool, error) {
if record == nil {
return false, nil
}
db = normalizeShardedDB(db)
for _, table := range observabilityShardTables("node_access_logs") {
var count int64
if err := db.Table(table).
Where(
"node_id = ? AND logged_at = ? AND remote_addr = ? AND host = ? AND path = ? AND status_code = ?",
record.NodeID,
record.LoggedAt,
record.RemoteAddr,
record.Host,
record.Path,
record.StatusCode,
).
Limit(1).
Count(&count).Error; err != nil {
return false, err
}
if count > 0 {
return true, nil
}
}
return false, nil
}
func DeleteNodeAccessLogsByNodeBefore(db *gorm.DB, nodeID string, before time.Time) (deleted int64, err error) {
return deleteAcrossShards(db, "node_access_logs", &NodeAccessLog{}, func(tx *gorm.DB) *gorm.DB {
return tx.Where("node_id = ? AND logged_at < ?", nodeID, before)
})
}
func applyNodeAccessLogFilters(db *gorm.DB, query NodeAccessLogQuery) *gorm.DB {
if trimmed := strings.TrimSpace(query.NodeID); trimmed != "" {
db = db.Where("node_id LIKE ?", "%"+trimmed+"%")
}
if trimmed := strings.TrimSpace(query.RemoteAddr); trimmed != "" {
db = db.Where("remote_addr LIKE ?", "%"+trimmed+"%")
}
if trimmed := strings.TrimSpace(query.Host); trimmed != "" {
db = db.Where("host LIKE ?", "%"+trimmed+"%")
}
if trimmed := strings.TrimSpace(query.Path); trimmed != "" {
db = db.Where("path LIKE ?", "%"+trimmed+"%")
}
if !query.Since.IsZero() {
db = db.Where("logged_at >= ?", query.Since)
}
if !query.Until.IsZero() {
db = db.Where("logged_at < ?", query.Until)
}
return db
}
func listNodeAccessLogsAcrossShards(query NodeAccessLogQuery) ([]*NodeAccessLog, error) {
items, err := queryAcrossShards("node_access_logs", func(tx *gorm.DB) ([]*NodeAccessLog, error) {
var shardRows []*NodeAccessLog
if err := applyNodeAccessLogFilters(tx, query).Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
sortNodeAccessLogs(items, query.SortBy, query.SortOrder)
return items, nil
}
func buildNodeAccessLogBucketRows(query NodeAccessLogBucketQuery) ([]*NodeAccessLogBucketRow, error) {
logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: query.RemoteAddr,
Host: query.Host,
Path: query.Path,
Since: query.Since,
})
if err != nil {
return nil, err
}
type bucketAccumulator struct {
requestCount int64
uniqueIPs map[string]struct{}
uniqueHosts map[string]struct{}
successCount int64
clientErrorCount int64
serverErrorCount int64
}
accumulators := make(map[int64]*bucketAccumulator)
for _, item := range logs {
if item == nil {
continue
}
bucketEpoch := bucketEpochForTime(item.LoggedAt, query.FoldMinutes)
accumulator := accumulators[bucketEpoch]
if accumulator == nil {
accumulator = &bucketAccumulator{
uniqueIPs: make(map[string]struct{}),
uniqueHosts: make(map[string]struct{}),
}
accumulators[bucketEpoch] = accumulator
}
accumulator.requestCount++
if trimmed := strings.TrimSpace(item.RemoteAddr); trimmed != "" {
accumulator.uniqueIPs[trimmed] = struct{}{}
}
if trimmed := strings.TrimSpace(item.Host); trimmed != "" {
accumulator.uniqueHosts[trimmed] = struct{}{}
}
switch {
case item.StatusCode < 400:
accumulator.successCount++
case item.StatusCode < 500:
accumulator.clientErrorCount++
default:
accumulator.serverErrorCount++
}
}
rows := make([]*NodeAccessLogBucketRow, 0, len(accumulators))
for bucketEpoch, accumulator := range accumulators {
rows = append(rows, &NodeAccessLogBucketRow{
BucketEpoch: bucketEpoch,
RequestCount: accumulator.requestCount,
UniqueIPCount: int64(len(accumulator.uniqueIPs)),
UniqueHostCount: int64(len(accumulator.uniqueHosts)),
SuccessCount: accumulator.successCount,
ClientErrorCount: accumulator.clientErrorCount,
ServerErrorCount: accumulator.serverErrorCount,
})
}
sortNodeAccessLogBucketRows(rows, query.SortBy, query.SortOrder)
return rows, nil
}
func buildNodeAccessLogBucketIPRows(query NodeAccessLogBucketIPQuery) ([]*NodeAccessLogBucketIPRow, error) {
if query.BucketStartedAt.IsZero() {
return []*NodeAccessLogBucketIPRow{}, nil
}
foldMinutes := query.FoldMinutes
if foldMinutes <= 0 {
foldMinutes = 3
}
bucketStartedAt := query.BucketStartedAt.UTC()
logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: query.RemoteAddr,
Host: query.Host,
Path: query.Path,
Since: bucketStartedAt,
Until: bucketStartedAt.Add(time.Duration(foldMinutes) * time.Minute),
})
if err != nil {
return nil, err
}
type accumulator struct {
requestCount int64
successCount int64
clientErrorCount int64
serverErrorCount int64
lastSeenAt time.Time
}
accumulators := make(map[string]*accumulator)
for _, item := range logs {
if item == nil {
continue
}
remoteAddr := strings.TrimSpace(item.RemoteAddr)
if remoteAddr == "" {
continue
}
acc := accumulators[remoteAddr]
if acc == nil {
acc = &accumulator{}
accumulators[remoteAddr] = acc
}
acc.requestCount++
switch {
case item.StatusCode < 400:
acc.successCount++
case item.StatusCode < 500:
acc.clientErrorCount++
default:
acc.serverErrorCount++
}
if item.LoggedAt.After(acc.lastSeenAt) {
acc.lastSeenAt = item.LoggedAt
}
}
rows := make([]*NodeAccessLogBucketIPRow, 0, len(accumulators))
for remoteAddr, acc := range accumulators {
rows = append(rows, &NodeAccessLogBucketIPRow{
RemoteAddr: remoteAddr,
RequestCount: acc.requestCount,
SuccessCount: acc.successCount,
ClientErrorCount: acc.clientErrorCount,
ServerErrorCount: acc.serverErrorCount,
LastSeenEpoch: acc.lastSeenAt.Unix(),
})
}
sortNodeAccessLogBucketIPRows(rows, query.SortBy, query.SortOrder)
return rows, nil
}
func buildNodeAccessLogIPSummaryRows(query NodeAccessLogIPSummaryQuery, recentSince time.Time) ([]*NodeAccessLogIPSummaryRow, error) {
logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: query.RemoteAddr,
Host: query.Host,
Since: query.Since,
})
if err != nil {
return nil, err
}
type accumulator struct {
totalRequests int64
recentRequests int64
lastSeenAt time.Time
}
accumulators := make(map[string]*accumulator)
for _, item := range logs {
if item == nil {
continue
}
remoteAddr := strings.TrimSpace(item.RemoteAddr)
if remoteAddr == "" {
continue
}
acc := accumulators[remoteAddr]
if acc == nil {
acc = &accumulator{}
accumulators[remoteAddr] = acc
}
acc.totalRequests++
if !recentSince.IsZero() && !item.LoggedAt.Before(recentSince) {
acc.recentRequests++
}
if item.LoggedAt.After(acc.lastSeenAt) {
acc.lastSeenAt = item.LoggedAt
}
}
rows := make([]*NodeAccessLogIPSummaryRow, 0, len(accumulators))
for remoteAddr, acc := range accumulators {
rows = append(rows, &NodeAccessLogIPSummaryRow{
RemoteAddr: remoteAddr,
TotalRequests: acc.totalRequests,
RecentRequests: acc.recentRequests,
LastSeenEpoch: acc.lastSeenAt.Unix(),
})
}
sortNodeAccessLogIPSummaryRows(rows, query.SortBy, query.SortOrder)
return rows, nil
}
func sortNodeAccessLogBucketIPRows(items []*NodeAccessLogBucketIPRow, sortBy string, sortOrder string) {
desc := normalizeSortOrder(sortOrder) != "asc"
sort.Slice(items, func(i int, j int) bool {
left := items[i]
right := items[j]
if left == nil || right == nil {
return left != nil
}
var compare int
switch strings.TrimSpace(sortBy) {
case "last_seen_at":
compare = compareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
case "remote_addr":
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
default:
compare = compareInt64(left.RequestCount, right.RequestCount)
}
if compare == 0 {
compare = compareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
}
if compare == 0 {
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
}
if desc {
return compare > 0
}
return compare < 0
})
}
func sortNodeAccessLogs(items []*NodeAccessLog, sortBy string, sortOrder string) {
desc := normalizeSortOrder(sortOrder) != "asc"
sort.Slice(items, func(i int, j int) bool {
left := items[i]
right := items[j]
if left == nil || right == nil {
return left != nil
}
var compare int
switch strings.TrimSpace(sortBy) {
case "status_code":
compare = compareInt(left.StatusCode, right.StatusCode)
case "remote_addr":
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
case "host":
compare = strings.Compare(left.Host, right.Host)
case "path":
compare = strings.Compare(left.Path, right.Path)
default:
compare = compareTime(left.LoggedAt, right.LoggedAt)
}
if compare == 0 {
compare = compareTime(left.LoggedAt, right.LoggedAt)
}
if compare == 0 {
compare = compareUint(left.ID, right.ID)
}
if desc {
return compare > 0
}
return compare < 0
})
}
func sortNodeAccessLogBucketRows(items []*NodeAccessLogBucketRow, sortBy string, sortOrder string) {
desc := normalizeSortOrder(sortOrder) != "asc"
sort.Slice(items, func(i int, j int) bool {
left := items[i]
right := items[j]
if left == nil || right == nil {
return left != nil
}
var compare int
switch strings.TrimSpace(sortBy) {
case "request_count":
compare = compareInt64(left.RequestCount, right.RequestCount)
default:
compare = compareInt64(left.BucketEpoch, right.BucketEpoch)
}
if compare == 0 {
compare = compareInt64(left.BucketEpoch, right.BucketEpoch)
}
if desc {
return compare > 0
}
return compare < 0
})
}
func sortNodeAccessLogIPSummaryRows(items []*NodeAccessLogIPSummaryRow, sortBy string, sortOrder string) {
desc := normalizeSortOrder(sortOrder) != "asc"
sort.Slice(items, func(i int, j int) bool {
left := items[i]
right := items[j]
if left == nil || right == nil {
return left != nil
}
var compare int
switch strings.TrimSpace(sortBy) {
case "recent_requests":
compare = compareInt64(left.RecentRequests, right.RecentRequests)
case "last_seen_at":
compare = compareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
case "remote_addr":
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
default:
compare = compareInt64(left.TotalRequests, right.TotalRequests)
}
if compare == 0 {
compare = compareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
}
if compare == 0 {
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
}
if desc {
return compare > 0
}
return compare < 0
})
}
func paginateBounds(total int, page int, pageSize int) (int, int) {
if page < 0 {
page = 0
}
if pageSize <= 0 {
return 0, total
}
start := page * pageSize
if start > total {
start = total
}
end := start + pageSize
if end > total {
end = total
}
return start, end
}
func bucketEpochForTime(value time.Time, bucketMinutes int) int64 {
bucketSeconds := int64(bucketMinutes * 60)
if bucketSeconds <= 0 {
bucketSeconds = 180
}
return (value.UTC().Unix() / bucketSeconds) * bucketSeconds
}
func compareTime(left time.Time, right time.Time) int {
switch {
case left.After(right):
return 1
case left.Before(right):
return -1
default:
return 0
}
}
func compareInt(left int, right int) int {
switch {
case left > right:
return 1
case left < right:
return -1
default:
return 0
}
}
func compareInt64(left int64, right int64) int {
switch {
case left > right:
return 1
case left < right:
return -1
default:
return 0
}
}
func compareUint(left uint, right uint) int {
switch {
case left > right:
return 1
case left < right:
return -1
default:
return 0
}
}
func normalizeSortOrder(sortOrder string) string {
if strings.EqualFold(strings.TrimSpace(sortOrder), "asc") {
return "asc"
}
return "desc"
}
@@ -0,0 +1,47 @@
package model
import "time"
type NodeHealthEvent struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
EventType string `json:"event_type" gorm:"index;size:64;not null"`
Severity string `json:"severity" gorm:"size:16;not null"`
Status string `json:"status" gorm:"index;size:16;not null"`
Message string `json:"message" gorm:"type:text"`
FirstTriggeredAt time.Time `json:"first_triggered_at" gorm:"index"`
LastTriggeredAt time.Time `json:"last_triggered_at" gorm:"index"`
ReportedAt time.Time `json:"reported_at" gorm:"index"`
ResolvedAt *time.Time `json:"resolved_at" gorm:"index"`
MetadataJSON string `json:"metadata_json" gorm:"type:text"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func GetActiveNodeHealthEvent(nodeID string, eventType string) (*NodeHealthEvent, error) {
event := &NodeHealthEvent{}
err := DB.Where("node_id = ? AND event_type = ? AND status = ?", nodeID, eventType, "active").First(event).Error
return event, err
}
func ListNodeHealthEvents(nodeID string, activeOnly bool, limit int) (events []*NodeHealthEvent, err error) {
query := DB.Where("node_id = ?", nodeID).Order("last_triggered_at desc")
if activeOnly {
query = query.Where("status = ?", "active")
}
if limit > 0 {
query = query.Limit(limit)
}
err = query.Find(&events).Error
return events, err
}
func ListActiveNodeHealthEvents() (events []*NodeHealthEvent, err error) {
err = DB.Where("status = ?", "active").Order("last_triggered_at desc").Find(&events).Error
return events, err
}
func DeleteNodeHealthEvents(nodeID string) (deleted int64, err error) {
result := DB.Where("node_id = ?", nodeID).Delete(&NodeHealthEvent{})
return result.RowsAffected, result.Error
}
@@ -0,0 +1,107 @@
package model
import (
"time"
"github.com/rain-kl/openflare/pkg/utils"
"gorm.io/gorm"
)
type NodeMetricSnapshot struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
CapturedAt time.Time `json:"captured_at" gorm:"index"`
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"`
NetworkRxBytes int64 `json:"network_rx_bytes"`
NetworkTxBytes int64 `json:"network_tx_bytes"`
CreatedAt time.Time `json:"created_at"`
}
func (snapshot *NodeMetricSnapshot) GetID() uint {
return snapshot.ID
}
func (snapshot *NodeMetricSnapshot) GetTime() time.Time {
return snapshot.CapturedAt
}
func (snapshot *NodeMetricSnapshot) BeforeCreate(tx *gorm.DB) error {
return assignObservabilityID(&snapshot.ID)
}
func (snapshot *NodeMetricSnapshot) Insert() error {
return DB.Create(snapshot).Error
}
func ListNodeMetricSnapshots(nodeID string, since time.Time, limit int) (snapshots []*NodeMetricSnapshot, err error) {
rows, err := queryAcrossShards("node_metric_snapshots", func(tx *gorm.DB) ([]*NodeMetricSnapshot, error) {
var shardRows []*NodeMetricSnapshot
query := tx.Order("captured_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("captured_at >= ?", since)
}
if err := query.Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
return utils.SortAndLimitRecords(rows, limit), nil
}
func ListMetricSnapshotsSince(since time.Time) (snapshots []*NodeMetricSnapshot, err error) {
rows, err := queryAcrossShards("node_metric_snapshots", func(tx *gorm.DB) ([]*NodeMetricSnapshot, error) {
var shardRows []*NodeMetricSnapshot
query := tx.Order("captured_at desc")
if !since.IsZero() {
query = query.Where("captured_at >= ?", since)
}
if err := query.Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
return utils.SortAndLimitRecords(rows, 0), nil
}
func NodeMetricSnapshotExists(db *gorm.DB, nodeID string, capturedAt time.Time) (bool, error) {
db = normalizeShardedDB(db)
for _, table := range observabilityShardTables("node_metric_snapshots") {
var count int64
if err := db.Table(table).
Where("node_id = ? AND captured_at = ?", nodeID, capturedAt).
Limit(1).
Count(&count).Error; err != nil {
return false, err
}
if count > 0 {
return true, nil
}
}
return false, nil
}
func DeleteNodeMetricSnapshotsBefore(db *gorm.DB, before time.Time) (int64, error) {
return deleteAcrossShards(db, "node_metric_snapshots", &NodeMetricSnapshot{}, func(tx *gorm.DB) *gorm.DB {
return tx.Where("captured_at < ?", before)
})
}
func DeleteAllNodeMetricSnapshots(db *gorm.DB) (int64, error) {
return deleteAcrossShards(db, "node_metric_snapshots", &NodeMetricSnapshot{}, nil)
}
@@ -0,0 +1,61 @@
package model
import (
"time"
"github.com/rain-kl/openflare/pkg/utils"
"gorm.io/gorm"
)
type NodeObservationFrpc struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
CapturedAt time.Time `json:"captured_at" gorm:"index"`
TunnelStatus string `json:"tunnel_status" gorm:"size:16"`
ConnectedRelaysCount int `json:"connected_relays_count"`
CreatedAt time.Time `json:"created_at"`
}
func (obs *NodeObservationFrpc) GetID() uint {
return obs.ID
}
func (obs *NodeObservationFrpc) GetTime() time.Time {
return obs.CapturedAt
}
func (obs *NodeObservationFrpc) BeforeCreate(tx *gorm.DB) error {
return assignObservabilityID(&obs.ID)
}
func (obs *NodeObservationFrpc) Insert() error {
return DB.Create(obs).Error
}
func ListNodeObservationFrpcs(nodeID string, since time.Time, limit int) (observations []*NodeObservationFrpc, err error) {
rows, err := queryAcrossShards("node_observation_frpcs", func(tx *gorm.DB) ([]*NodeObservationFrpc, error) {
var shardRows []*NodeObservationFrpc
query := tx.Order("captured_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("captured_at >= ?", since)
}
if err := query.Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
return utils.SortAndLimitRecords(rows, limit), nil
}
func DeleteNodeObservationFrpcsBefore(db *gorm.DB, before time.Time) (int64, error) {
return deleteAcrossShards(db, "node_observation_frpcs", &NodeObservationFrpc{}, func(tx *gorm.DB) *gorm.DB {
return tx.Where("captured_at < ?", before)
})
}
@@ -0,0 +1,63 @@
package model
import (
"time"
"github.com/rain-kl/openflare/pkg/utils"
"gorm.io/gorm"
)
type NodeObservationFrps struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
CapturedAt time.Time `json:"captured_at" gorm:"index"`
FrpsConnections int `json:"frps_connections"`
FrpsProxyCount int `json:"frps_proxy_count"`
FrpsClientCount int `json:"frps_client_count"`
FrpsProxies string `json:"frps_proxies" gorm:"type:text"`
CreatedAt time.Time `json:"created_at"`
}
func (obs *NodeObservationFrps) GetID() uint {
return obs.ID
}
func (obs *NodeObservationFrps) GetTime() time.Time {
return obs.CapturedAt
}
func (obs *NodeObservationFrps) BeforeCreate(tx *gorm.DB) error {
return assignObservabilityID(&obs.ID)
}
func (obs *NodeObservationFrps) Insert() error {
return DB.Create(obs).Error
}
func ListNodeObservationFrps(nodeID string, since time.Time, limit int) (observations []*NodeObservationFrps, err error) {
rows, err := queryAcrossShards("node_observation_frps", func(tx *gorm.DB) ([]*NodeObservationFrps, error) {
var shardRows []*NodeObservationFrps
query := tx.Order("captured_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("captured_at >= ?", since)
}
if err := query.Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
return utils.SortAndLimitRecords(rows, limit), nil
}
func DeleteNodeObservationFrpsBefore(db *gorm.DB, before time.Time) (int64, error) {
return deleteAcrossShards(db, "node_observation_frps", &NodeObservationFrps{}, func(tx *gorm.DB) *gorm.DB {
return tx.Where("captured_at < ?", before)
})
}
@@ -0,0 +1,62 @@
package model
import (
"time"
"github.com/rain-kl/openflare/pkg/utils"
"gorm.io/gorm"
)
type NodeObservationOpenresty struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
CapturedAt time.Time `json:"captured_at" gorm:"index"`
OpenrestyRxBytes int64 `json:"openresty_rx_bytes"`
OpenrestyTxBytes int64 `json:"openresty_tx_bytes"`
OpenrestyConnections int64 `json:"openresty_connections"`
CreatedAt time.Time `json:"created_at"`
}
func (obs *NodeObservationOpenresty) GetID() uint {
return obs.ID
}
func (obs *NodeObservationOpenresty) GetTime() time.Time {
return obs.CapturedAt
}
func (obs *NodeObservationOpenresty) BeforeCreate(tx *gorm.DB) error {
return assignObservabilityID(&obs.ID)
}
func (obs *NodeObservationOpenresty) Insert() error {
return DB.Create(obs).Error
}
func ListNodeObservationOpenresty(nodeID string, since time.Time, limit int) (observations []*NodeObservationOpenresty, err error) {
rows, err := queryAcrossShards("node_observation_openresties", func(tx *gorm.DB) ([]*NodeObservationOpenresty, error) {
var shardRows []*NodeObservationOpenresty
query := tx.Order("captured_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("captured_at >= ?", since)
}
if err := query.Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
return utils.SortAndLimitRecords(rows, limit), nil
}
func DeleteNodeObservationOpenrestiesBefore(db *gorm.DB, before time.Time) (int64, error) {
return deleteAcrossShards(db, "node_observation_openresties", &NodeObservationOpenresty{}, func(tx *gorm.DB) *gorm.DB {
return tx.Where("captured_at < ?", before)
})
}
@@ -0,0 +1,105 @@
package model
import (
"time"
"github.com/rain-kl/openflare/pkg/utils"
"gorm.io/gorm"
)
type NodeRequestReport struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
WindowStartedAt time.Time `json:"window_started_at" gorm:"index"`
WindowEndedAt time.Time `json:"window_ended_at" gorm:"index"`
RequestCount int64 `json:"request_count"`
ErrorCount int64 `json:"error_count"`
UniqueVisitorCount int64 `json:"unique_visitor_count"`
StatusCodesJSON string `json:"status_codes_json" gorm:"type:text"`
TopDomainsJSON string `json:"top_domains_json" gorm:"type:text"`
SourceCountriesJSON string `json:"source_countries_json" gorm:"type:text"`
CreatedAt time.Time `json:"created_at"`
}
func (report *NodeRequestReport) GetID() uint {
return report.ID
}
func (report *NodeRequestReport) GetTime() time.Time {
return report.WindowEndedAt
}
func (report *NodeRequestReport) BeforeCreate(tx *gorm.DB) error {
return assignObservabilityID(&report.ID)
}
func (report *NodeRequestReport) Insert() error {
return DB.Create(report).Error
}
func ListNodeRequestReports(nodeID string, since time.Time, limit int) (reports []*NodeRequestReport, err error) {
rows, err := queryAcrossShards("node_request_reports", func(tx *gorm.DB) ([]*NodeRequestReport, error) {
var shardRows []*NodeRequestReport
query := tx.Order("window_ended_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("window_ended_at >= ?", since)
}
if err := query.Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
return utils.SortAndLimitRecords(rows, limit), nil
}
func ListRequestReportsSince(since time.Time) (reports []*NodeRequestReport, err error) {
rows, err := queryAcrossShards("node_request_reports", func(tx *gorm.DB) ([]*NodeRequestReport, error) {
var shardRows []*NodeRequestReport
query := tx.Order("window_ended_at desc")
if !since.IsZero() {
query = query.Where("window_ended_at >= ?", since)
}
if err := query.Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
return utils.SortAndLimitRecords(rows, 0), nil
}
func NodeRequestReportExists(db *gorm.DB, nodeID string, windowStartedAt time.Time, windowEndedAt time.Time) (bool, error) {
db = normalizeShardedDB(db)
for _, table := range observabilityShardTables("node_request_reports") {
var count int64
if err := db.Table(table).
Where("node_id = ? AND window_started_at = ? AND window_ended_at = ?", nodeID, windowStartedAt, windowEndedAt).
Limit(1).
Count(&count).Error; err != nil {
return false, err
}
if count > 0 {
return true, nil
}
}
return false, nil
}
func DeleteNodeRequestReportsBefore(db *gorm.DB, before time.Time) (int64, error) {
return deleteAcrossShards(db, "node_request_reports", &NodeRequestReport{}, func(tx *gorm.DB) *gorm.DB {
return tx.Where("window_ended_at < ?", before)
})
}
func DeleteAllNodeRequestReports(db *gorm.DB) (int64, error) {
return deleteAcrossShards(db, "node_request_reports", &NodeRequestReport{}, nil)
}
@@ -0,0 +1,54 @@
package model
import (
"time"
"gorm.io/gorm/clause"
)
type NodeSystemProfile struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"uniqueIndex;size:64;not null"`
Hostname string `json:"hostname" gorm:"size:255"`
OSName string `json:"os_name" gorm:"size:128"`
OSVersion string `json:"os_version" gorm:"size:128"`
KernelVersion string `json:"kernel_version" gorm:"size:128"`
Architecture string `json:"architecture" gorm:"size:64"`
CPUModel string `json:"cpu_model" gorm:"size:255"`
CPUCores int `json:"cpu_cores"`
TotalMemoryBytes int64 `json:"total_memory_bytes"`
TotalDiskBytes int64 `json:"total_disk_bytes"`
UptimeSeconds int64 `json:"uptime_seconds"`
ReportedAt time.Time `json:"reported_at" gorm:"index"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func GetNodeSystemProfile(nodeID string) (*NodeSystemProfile, error) {
profile := &NodeSystemProfile{}
err := DB.Where("node_id = ?", nodeID).First(profile).Error
return profile, err
}
func UpsertNodeSystemProfile(profile *NodeSystemProfile) error {
if profile == nil {
return nil
}
return DB.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "node_id"}},
DoUpdates: clause.AssignmentColumns([]string{
"hostname",
"os_name",
"os_version",
"kernel_version",
"architecture",
"cpu_model",
"cpu_cores",
"total_memory_bytes",
"total_disk_bytes",
"uptime_seconds",
"reported_at",
"updated_at",
}),
}).Create(profile).Error
}
+426
View File
@@ -0,0 +1,426 @@
package model
import (
"strconv"
"strings"
"time"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/pkg/geoip"
"gorm.io/gorm"
)
type Option struct {
Key string `json:"key" gorm:"primaryKey"`
Value string `json:"value"`
}
func AllOption() ([]*Option, error) {
var options []*Option
var err error
err = DB.Find(&options).Error
return options, err
}
func InitOptionMap() {
common.OptionMapRWMutex.Lock()
common.OptionMap = make(map[string]string)
common.OptionMap["PasswordLoginEnabled"] = strconv.FormatBool(common.PasswordLoginEnabled)
common.OptionMap["CapLoginEnabled"] = strconv.FormatBool(common.CapLoginEnabled)
common.OptionMap["PasswordRegisterEnabled"] = strconv.FormatBool(common.PasswordRegisterEnabled)
common.OptionMap["EmailVerificationEnabled"] = strconv.FormatBool(common.EmailVerificationEnabled)
common.OptionMap["GitHubOAuthEnabled"] = strconv.FormatBool(common.GitHubOAuthEnabled)
common.OptionMap["WeChatAuthEnabled"] = strconv.FormatBool(common.WeChatAuthEnabled)
common.OptionMap["SMTPServer"] = ""
common.OptionMap["SMTPPort"] = strconv.Itoa(common.SMTPPort)
common.OptionMap["SMTPAccount"] = ""
common.OptionMap["SMTPToken"] = ""
common.OptionMap["Notice"] = ""
common.OptionMap["About"] = ""
common.OptionMap["Footer"] = common.Footer
common.OptionMap["HomePageLink"] = common.HomePageLink
common.OptionMap["SystemName"] = common.SystemName
common.OptionMap["ServerAddress"] = ""
common.OptionMap["GitHubClientId"] = ""
common.OptionMap["GitHubClientSecret"] = ""
common.OptionMap["WeChatServerAddress"] = ""
common.OptionMap["WeChatServerToken"] = ""
common.OptionMap["WeChatAccountQRCodeImageURL"] = ""
common.OptionMap["AgentDiscoveryToken"] = ""
common.OptionMap["AgentHeartbeatInterval"] = strconv.Itoa(common.AgentHeartbeatInterval)
common.OptionMap["AgentWebsocketUpgradeEnabled"] = strconv.FormatBool(common.AgentWebsocketUpgradeEnabled)
common.OptionMap["NodeOfflineThreshold"] = strconv.Itoa(int(common.NodeOfflineThreshold.Milliseconds()))
common.OptionMap["AgentUpdateRepo"] = common.AgentUpdateRepo
common.OptionMap["GeoIPProvider"] = common.GeoIPProvider
common.OptionMap["DatabaseAutoCleanupEnabled"] = strconv.FormatBool(common.DatabaseAutoCleanupEnabled)
common.OptionMap["UptimeKumaEnabled"] = strconv.FormatBool(common.UptimeKumaEnabled)
common.OptionMap["UptimeKumaUrl"] = common.UptimeKumaUrl
common.OptionMap["UptimeKumaUsername"] = common.UptimeKumaUsername
common.OptionMap["UptimeKumaPassword"] = common.UptimeKumaPassword
common.OptionMap["UptimeKumaMonitorScope"] = common.UptimeKumaMonitorScope
common.OptionMap["UptimeKumaSelectedSites"] = common.UptimeKumaSelectedSites
common.OptionMap["UptimeKumaSyncInterval"] = strconv.Itoa(common.UptimeKumaSyncInterval)
common.OptionMap["UptimeKumaInterval"] = strconv.Itoa(common.UptimeKumaInterval)
common.OptionMap["UptimeKumaRetry"] = strconv.Itoa(common.UptimeKumaRetry)
common.OptionMap["UptimeKumaRetryInterval"] = strconv.Itoa(common.UptimeKumaRetryInterval)
common.OptionMap["UptimeKumaTimeout"] = strconv.Itoa(common.UptimeKumaTimeout)
common.OptionMap["DatabaseAutoCleanupRetentionDays"] = strconv.Itoa(common.DatabaseAutoCleanupRetentionDays)
common.OptionMap["OpenRestyDefaultServerReturnStatus"] = strconv.Itoa(common.OpenRestyDefaultServerReturnStatus)
common.OptionMap["OpenRestyWorkerProcesses"] = common.OpenRestyWorkerProcesses
common.OptionMap["OpenRestyWorkerConnections"] = strconv.Itoa(common.OpenRestyWorkerConnections)
common.OptionMap["OpenRestyWorkerRlimitNofile"] = strconv.Itoa(common.OpenRestyWorkerRlimitNofile)
common.OptionMap["OpenRestyEventsUse"] = common.OpenRestyEventsUse
common.OptionMap["OpenRestyEventsMultiAcceptEnabled"] = strconv.FormatBool(common.OpenRestyEventsMultiAcceptEnabled)
common.OptionMap["OpenRestyKeepaliveTimeout"] = strconv.Itoa(common.OpenRestyKeepaliveTimeout)
common.OptionMap["OpenRestyKeepaliveRequests"] = strconv.Itoa(common.OpenRestyKeepaliveRequests)
common.OptionMap["OpenRestyClientHeaderTimeout"] = strconv.Itoa(common.OpenRestyClientHeaderTimeout)
common.OptionMap["OpenRestyClientBodyTimeout"] = strconv.Itoa(common.OpenRestyClientBodyTimeout)
common.OptionMap["OpenRestyClientMaxBodySize"] = common.OpenRestyClientMaxBodySize
common.OptionMap["OpenRestyLargeClientHeaderBuffers"] = common.OpenRestyLargeClientHeaderBuffers
common.OptionMap["OpenRestySendTimeout"] = strconv.Itoa(common.OpenRestySendTimeout)
common.OptionMap["OpenRestyProxyConnectTimeout"] = strconv.Itoa(common.OpenRestyProxyConnectTimeout)
common.OptionMap["OpenRestyProxySendTimeout"] = strconv.Itoa(common.OpenRestyProxySendTimeout)
common.OptionMap["OpenRestyProxyReadTimeout"] = strconv.Itoa(common.OpenRestyProxyReadTimeout)
common.OptionMap["OpenRestyWebsocketEnabled"] = strconv.FormatBool(common.OpenRestyWebsocketEnabled)
common.OptionMap["OpenRestyHTTP3Enabled"] = strconv.FormatBool(common.OpenRestyHTTP3Enabled)
common.OptionMap["OpenRestyProxyRequestBufferingEnabled"] = strconv.FormatBool(common.OpenRestyProxyRequestBufferingEnabled)
common.OptionMap["OpenRestyProxyBufferingEnabled"] = strconv.FormatBool(common.OpenRestyProxyBufferingEnabled)
common.OptionMap["OpenRestyProxyBuffers"] = common.OpenRestyProxyBuffers
common.OptionMap["OpenRestyProxyBufferSize"] = common.OpenRestyProxyBufferSize
common.OptionMap["OpenRestyProxyBusyBuffersSize"] = common.OpenRestyProxyBusyBuffersSize
common.OptionMap["OpenRestyGzipEnabled"] = strconv.FormatBool(common.OpenRestyGzipEnabled)
common.OptionMap["OpenRestyGzipMinLength"] = strconv.Itoa(common.OpenRestyGzipMinLength)
common.OptionMap["OpenRestyGzipCompLevel"] = strconv.Itoa(common.OpenRestyGzipCompLevel)
common.OptionMap["OpenRestyCacheEnabled"] = strconv.FormatBool(common.OpenRestyCacheEnabled)
common.OptionMap["OpenRestyCachePath"] = common.OpenRestyCachePath
common.OptionMap["OpenRestyCacheLevels"] = common.OpenRestyCacheLevels
common.OptionMap["OpenRestyCacheInactive"] = common.OpenRestyCacheInactive
common.OptionMap["OpenRestyCacheMaxSize"] = common.OpenRestyCacheMaxSize
common.OptionMap["OpenRestyCacheKeyTemplate"] = common.OpenRestyCacheKeyTemplate
common.OptionMap["OpenRestyCacheLockEnabled"] = strconv.FormatBool(common.OpenRestyCacheLockEnabled)
common.OptionMap["OpenRestyCacheLockTimeout"] = common.OpenRestyCacheLockTimeout
common.OptionMap["OpenRestyCacheUseStale"] = common.OpenRestyCacheUseStale
common.OptionMap["OpenRestyMainConfigTemplate"] = common.OpenRestyMainConfigTemplate
common.OptionMap["GlobalApiRateLimitNum"] = strconv.Itoa(common.GlobalApiRateLimitNum)
common.OptionMap["GlobalApiRateLimitDuration"] = strconv.FormatInt(common.GlobalApiRateLimitDuration, 10)
common.OptionMap["GlobalWebRateLimitNum"] = strconv.Itoa(common.GlobalWebRateLimitNum)
common.OptionMap["GlobalWebRateLimitDuration"] = strconv.FormatInt(common.GlobalWebRateLimitDuration, 10)
common.OptionMap["CriticalRateLimitNum"] = strconv.Itoa(common.CriticalRateLimitNum)
common.OptionMap["CriticalRateLimitDuration"] = strconv.FormatInt(common.CriticalRateLimitDuration, 10)
common.OptionMapRWMutex.Unlock()
options, _ := AllOption()
for _, option := range options {
updateOptionMap(option.Key, option.Value)
}
}
func UpdateOption(key string, value string) error {
return UpdateOptions([]Option{{
Key: key,
Value: value,
}})
}
func UpdateOptions(options []Option) error {
if len(options) == 0 {
return nil
}
if err := DB.Transaction(func(tx *gorm.DB) error {
for _, item := range options {
if item.Key == "UptimeKumaPassword" && strings.TrimSpace(item.Value) == "" {
continue
}
option := Option{
Key: item.Key,
}
if err := tx.FirstOrCreate(&option, Option{Key: item.Key}).Error; err != nil {
return err
}
option.Value = item.Value
if err := tx.Save(&option).Error; err != nil {
return err
}
}
return nil
}); err != nil {
return err
}
for _, item := range options {
if item.Key == "UptimeKumaPassword" && strings.TrimSpace(item.Value) == "" {
continue
}
updateOptionMap(item.Key, item.Value)
}
return nil
}
func updateOptionMap(key string, value string) {
shouldRefreshGeoIP := false
common.OptionMapRWMutex.Lock()
if common.OptionMap == nil {
common.OptionMap = make(map[string]string)
}
common.OptionMap[key] = value
if strings.HasSuffix(key, "Enabled") {
boolValue := value == "true"
switch key {
case "PasswordRegisterEnabled":
common.PasswordRegisterEnabled = boolValue
case "PasswordLoginEnabled":
common.PasswordLoginEnabled = boolValue
case "CapLoginEnabled":
common.CapLoginEnabled = boolValue
case "EmailVerificationEnabled":
common.EmailVerificationEnabled = boolValue
case "GitHubOAuthEnabled":
common.GitHubOAuthEnabled = boolValue
case "WeChatAuthEnabled":
common.WeChatAuthEnabled = boolValue
}
}
switch key {
case "SMTPServer":
common.SMTPServer = value
case "SMTPPort":
intValue, _ := strconv.Atoi(value)
common.SMTPPort = intValue
case "SMTPAccount":
common.SMTPAccount = value
case "SMTPToken":
common.SMTPToken = value
case "ServerAddress":
common.ServerAddress = value
case "GitHubClientId":
common.GitHubClientId = value
case "GitHubClientSecret":
common.GitHubClientSecret = value
case "Footer":
common.Footer = value
case "HomePageLink":
common.HomePageLink = value
case "SystemName":
common.SystemName = value
case "WeChatServerAddress":
common.WeChatServerAddress = value
case "WeChatServerToken":
common.WeChatServerToken = value
case "WeChatAccountQRCodeImageURL":
common.WeChatAccountQRCodeImageURL = value
case "AgentDiscoveryToken":
common.AgentDiscoveryToken = value
case "AgentHeartbeatInterval":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.AgentHeartbeatInterval = v
}
case "AgentWebsocketUpgradeEnabled":
common.AgentWebsocketUpgradeEnabled = value == "true"
case "NodeOfflineThreshold":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.NodeOfflineThreshold = time.Duration(v) * time.Millisecond
}
case "AgentUpdateRepo":
if value != "" {
common.AgentUpdateRepo = value
}
case "GeoIPProvider":
if geoip.IsValidProvider(value) {
common.GeoIPProvider = value
shouldRefreshGeoIP = true
}
case "UptimeKumaEnabled":
common.UptimeKumaEnabled = value == "true"
case "UptimeKumaUrl":
common.UptimeKumaUrl = value
case "UptimeKumaUsername":
common.UptimeKumaUsername = value
case "UptimeKumaPassword":
common.UptimeKumaPassword = value
case "UptimeKumaMonitorScope":
common.UptimeKumaMonitorScope = value
case "UptimeKumaSelectedSites":
common.UptimeKumaSelectedSites = value
case "UptimeKumaSyncInterval":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.UptimeKumaSyncInterval = v
}
case "UptimeKumaInterval":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.UptimeKumaInterval = v
}
case "UptimeKumaRetry":
if v, err := strconv.Atoi(value); err == nil && v >= 0 {
common.UptimeKumaRetry = v
}
case "UptimeKumaRetryInterval":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.UptimeKumaRetryInterval = v
}
case "UptimeKumaTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.UptimeKumaTimeout = v
}
case "DatabaseAutoCleanupEnabled":
common.DatabaseAutoCleanupEnabled = value == "true"
case "DatabaseAutoCleanupRetentionDays":
if v, err := strconv.Atoi(value); err == nil && v >= 1 {
common.DatabaseAutoCleanupRetentionDays = v
}
case "OpenRestyDefaultServerReturnStatus":
if v, err := strconv.Atoi(value); err == nil && v >= 100 && v <= 999 {
common.OpenRestyDefaultServerReturnStatus = v
}
case "OpenRestyWorkerProcesses":
if strings.TrimSpace(value) != "" {
common.OpenRestyWorkerProcesses = value
}
case "OpenRestyWorkerConnections":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyWorkerConnections = v
}
case "OpenRestyWorkerRlimitNofile":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyWorkerRlimitNofile = v
}
case "OpenRestyEventsUse":
common.OpenRestyEventsUse = value
case "OpenRestyResolvers":
common.OpenRestyResolvers = value
case "OpenRestyEventsMultiAcceptEnabled":
common.OpenRestyEventsMultiAcceptEnabled = value == "true"
case "OpenRestyKeepaliveTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyKeepaliveTimeout = v
}
case "OpenRestyKeepaliveRequests":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyKeepaliveRequests = v
}
case "OpenRestyClientHeaderTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyClientHeaderTimeout = v
}
case "OpenRestyClientBodyTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyClientBodyTimeout = v
}
case "OpenRestyClientMaxBodySize":
if strings.TrimSpace(value) != "" {
common.OpenRestyClientMaxBodySize = value
}
case "OpenRestyLargeClientHeaderBuffers":
if strings.TrimSpace(value) != "" {
common.OpenRestyLargeClientHeaderBuffers = value
}
case "OpenRestySendTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestySendTimeout = v
}
case "OpenRestyProxyConnectTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyProxyConnectTimeout = v
}
case "OpenRestyProxySendTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyProxySendTimeout = v
}
case "OpenRestyProxyReadTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyProxyReadTimeout = v
}
case "OpenRestyWebsocketEnabled":
common.OpenRestyWebsocketEnabled = value == "true"
case "OpenRestyHTTP3Enabled":
common.OpenRestyHTTP3Enabled = value == "true"
case "OpenRestyProxyRequestBufferingEnabled":
common.OpenRestyProxyRequestBufferingEnabled = value == "true"
case "OpenRestyProxyBufferingEnabled":
common.OpenRestyProxyBufferingEnabled = value == "true"
case "OpenRestyProxyBuffers":
if strings.TrimSpace(value) != "" {
common.OpenRestyProxyBuffers = value
}
case "OpenRestyProxyBufferSize":
if strings.TrimSpace(value) != "" {
common.OpenRestyProxyBufferSize = value
}
case "OpenRestyProxyBusyBuffersSize":
if strings.TrimSpace(value) != "" {
common.OpenRestyProxyBusyBuffersSize = value
}
case "OpenRestyGzipEnabled":
common.OpenRestyGzipEnabled = value == "true"
case "OpenRestyGzipMinLength":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyGzipMinLength = v
}
case "OpenRestyGzipCompLevel":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyGzipCompLevel = v
}
case "OpenRestyCacheEnabled":
common.OpenRestyCacheEnabled = value == "true"
case "OpenRestyCachePath":
common.OpenRestyCachePath = value
case "OpenRestyCacheLevels":
if strings.TrimSpace(value) != "" {
common.OpenRestyCacheLevels = value
}
case "OpenRestyCacheInactive":
if strings.TrimSpace(value) != "" {
common.OpenRestyCacheInactive = value
}
case "OpenRestyCacheMaxSize":
if strings.TrimSpace(value) != "" {
common.OpenRestyCacheMaxSize = value
}
case "OpenRestyCacheKeyTemplate":
if strings.TrimSpace(value) != "" {
common.OpenRestyCacheKeyTemplate = value
}
case "OpenRestyCacheLockEnabled":
common.OpenRestyCacheLockEnabled = value == "true"
case "OpenRestyCacheLockTimeout":
if strings.TrimSpace(value) != "" {
common.OpenRestyCacheLockTimeout = value
}
case "OpenRestyCacheUseStale":
if strings.TrimSpace(value) != "" {
common.OpenRestyCacheUseStale = value
}
case "OpenRestyMainConfigTemplate":
if strings.TrimSpace(value) != "" {
common.OpenRestyMainConfigTemplate = value
}
case "GlobalApiRateLimitNum":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.GlobalApiRateLimitNum = v
}
case "GlobalApiRateLimitDuration":
if v, err := strconv.ParseInt(value, 10, 64); err == nil && v > 0 {
common.GlobalApiRateLimitDuration = v
}
case "GlobalWebRateLimitNum":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.GlobalWebRateLimitNum = v
}
case "GlobalWebRateLimitDuration":
if v, err := strconv.ParseInt(value, 10, 64); err == nil && v > 0 {
common.GlobalWebRateLimitDuration = v
}
case "CriticalRateLimitNum":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.CriticalRateLimitNum = v
}
case "CriticalRateLimitDuration":
if v, err := strconv.ParseInt(value, 10, 64); err == nil && v > 0 {
common.CriticalRateLimitDuration = v
}
}
common.OptionMapRWMutex.Unlock()
if shouldRefreshGeoIP {
geoip.InitGeoIP(common.GeoIPProvider)
}
}
+56
View File
@@ -0,0 +1,56 @@
package model
import "time"
type Origin struct {
ID uint `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"size:255;not null"`
Address string `json:"address" gorm:"uniqueIndex;size:255;not null"`
Remark string `json:"remark" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type OriginRouteCount struct {
OriginID uint `json:"origin_id"`
RouteCount int64 `json:"route_count"`
}
func ListOrigins() (origins []*Origin, err error) {
err = DB.Order("id desc").Find(&origins).Error
return origins, err
}
func GetOriginByID(id uint) (*Origin, error) {
origin := &Origin{}
err := DB.First(origin, id).Error
return origin, err
}
func GetOriginByAddress(address string) (*Origin, error) {
origin := &Origin{}
err := DB.Where("address = ?", address).First(origin).Error
return origin, err
}
func ListOriginRouteCounts() ([]OriginRouteCount, error) {
result := make([]OriginRouteCount, 0)
err := DB.Model(&ProxyRoute{}).
Select("origin_id, COUNT(*) AS route_count").
Where("origin_id IS NOT NULL").
Group("origin_id").
Scan(&result).Error
return result, err
}
func (origin *Origin) Insert() error {
return DB.Create(origin).Error
}
func (origin *Origin) Update() error {
return DB.Save(origin).Error
}
func (origin *Origin) Delete() error {
return DB.Delete(origin).Error
}
+83
View File
@@ -0,0 +1,83 @@
package model
import "time"
const (
PagesDeploymentStatusUploaded = "uploaded"
PagesDeploymentStatusActive = "active"
)
type PagesProject struct {
ID uint `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"size:255;not null"`
Slug string `json:"slug" gorm:"uniqueIndex;size:128;not null"`
Description string `json:"description" gorm:"type:text;not null;default:''"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
SPAFallbackEnabled bool `json:"spa_fallback_enabled" gorm:"not null;default:false"`
SPAFallbackPath string `json:"spa_fallback_path" gorm:"size:512;not null;default:'/index.html'"`
APIProxyEnabled bool `json:"api_proxy_enabled" gorm:"not null;default:false"`
APIProxyPath string `json:"api_proxy_path" gorm:"size:255;not null;default:''"`
APIProxyPass string `json:"api_proxy_pass" gorm:"size:2048;not null;default:''"`
APIProxyRewrite string `json:"api_proxy_rewrite" gorm:"size:255;not null;default:''"`
ActiveDeploymentID *uint `json:"active_deployment_id" gorm:"index"`
RootDir string `json:"root_dir" gorm:"size:512;not null;default:''"`
EntryFile string `json:"entry_file" gorm:"size:512;not null;default:'index.html'"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type PagesDeployment struct {
ID uint `json:"id" gorm:"primaryKey"`
ProjectID uint `json:"project_id" gorm:"not null;index"`
DeploymentNumber int `json:"deployment_number" gorm:"not null"`
Checksum string `json:"checksum" gorm:"size:64;not null;index"`
Status string `json:"status" gorm:"size:32;not null;default:'uploaded';index"`
ArtifactPath string `json:"artifact_path" gorm:"size:2048;not null"`
FileCount int `json:"file_count" gorm:"not null;default:0"`
TotalSize int64 `json:"total_size" gorm:"not null;default:0"`
CreatedBy string `json:"created_by" gorm:"size:64;not null;default:''"`
CreatedAt time.Time `json:"created_at"`
ActivatedAt *time.Time `json:"activated_at"`
}
type PagesDeploymentFile struct {
ID uint `json:"id" gorm:"primaryKey"`
DeploymentID uint `json:"deployment_id" gorm:"not null;index"`
Path string `json:"path" gorm:"size:2048;not null"`
Size int64 `json:"size" gorm:"not null;default:0"`
Checksum string `json:"checksum" gorm:"size:64;not null"`
CreatedAt time.Time `json:"created_at"`
}
func ListPagesProjects() (projects []*PagesProject, err error) {
err = DB.Order("id desc").Find(&projects).Error
return projects, err
}
func GetPagesProjectByID(id uint) (*PagesProject, error) {
project := &PagesProject{}
err := DB.First(project, id).Error
return project, err
}
func GetPagesProjectBySlug(slug string) (*PagesProject, error) {
project := &PagesProject{}
err := DB.Where("slug = ?", slug).First(project).Error
return project, err
}
func ListPagesDeployments(projectID uint) (deployments []*PagesDeployment, err error) {
err = DB.Where("project_id = ?", projectID).Order("id desc").Find(&deployments).Error
return deployments, err
}
func GetPagesDeploymentByID(id uint) (*PagesDeployment, error) {
deployment := &PagesDeployment{}
err := DB.First(deployment, id).Error
return deployment, err
}
func ListPagesDeploymentFiles(deploymentID uint) (files []*PagesDeploymentFile, err error) {
err = DB.Where("deployment_id = ?", deploymentID).Order("path asc").Find(&files).Error
return files, err
}
@@ -0,0 +1,101 @@
package model
import "time"
type ProxyRoute struct {
ID uint `json:"id" gorm:"primaryKey"`
SiteName string `json:"site_name" gorm:"size:255;not null;default:''"`
Domain string `json:"domain" gorm:"uniqueIndex;size:255;not null"`
Domains string `json:"domains" gorm:"type:text;not null;default:'[]'"`
OriginID *uint `json:"origin_id" gorm:"index"`
OriginURL string `json:"origin_url" gorm:"size:2048;not null"`
OriginHost string `json:"origin_host" gorm:"size:255"`
Upstreams string `json:"upstreams" gorm:"type:text;not null;default:'[]'"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
EnableHTTPS bool `json:"enable_https" gorm:"column:enable_https;not null;default:false"`
CertID *uint `json:"cert_id"`
CertIDs string `json:"cert_ids" gorm:"type:text;not null;default:'[]'"`
DomainCertIDs string `json:"domain_cert_ids" gorm:"type:text;not null;default:'[]'"`
RedirectHTTP bool `json:"redirect_http" gorm:"not null;default:false"`
LimitConnPerServer int `json:"limit_conn_per_server" gorm:"not null;default:0"`
LimitConnPerIP int `json:"limit_conn_per_ip" gorm:"not null;default:0"`
LimitRate string `json:"limit_rate" gorm:"size:32;not null;default:''"`
CacheEnabled bool `json:"cache_enabled" gorm:"not null;default:false"`
CachePolicy string `json:"cache_policy" gorm:"size:32;not null;default:''"`
CacheRules string `json:"cache_rules" gorm:"type:text;not null;default:'[]'"`
CustomHeaders string `json:"custom_headers" gorm:"type:text;not null;default:'[]'"`
BasicAuthEnabled bool `json:"basic_auth_enabled" gorm:"not null;default:false"`
BasicAuthUsername string `json:"basic_auth_username" gorm:"size:255;not null;default:''"`
BasicAuthPassword string `json:"basic_auth_password" gorm:"size:255;not null;default:''"`
Remark string `json:"remark" gorm:"size:255"`
UpstreamType string `json:"upstream_type" gorm:"size:32;not null;default:'direct'"`
TunnelNodeID *uint `json:"tunnel_node_id" gorm:"index"`
TunnelTargetAddr string `json:"tunnel_target_addr" gorm:"size:512"`
TunnelTargetProtocol string `json:"tunnel_target_protocol" gorm:"size:16"`
PagesProjectID *uint `json:"pages_project_id" gorm:"index"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func ListProxyRoutes() (routes []*ProxyRoute, err error) {
err = DB.Order("id desc").Find(&routes).Error
return routes, err
}
func GetEnabledProxyRoutes() (routes []*ProxyRoute, err error) {
err = DB.Where("enabled = ?", true).Order("site_name asc").Order("domain asc").Find(&routes).Error
return routes, err
}
func GetProxyRouteByID(id uint) (*ProxyRoute, error) {
route := &ProxyRoute{}
err := DB.First(route, id).Error
return route, err
}
func ListProxyRoutesByOriginID(originID uint) (routes []*ProxyRoute, err error) {
err = DB.Where("origin_id = ?", originID).Order("id desc").Find(&routes).Error
return routes, err
}
func (route *ProxyRoute) Insert() error {
return DB.Create(route).Error
}
func (route *ProxyRoute) Update() error {
return DB.Model(&ProxyRoute{}).Where("id = ?", route.ID).Updates(map[string]any{
"site_name": route.SiteName,
"domain": route.Domain,
"domains": route.Domains,
"origin_id": route.OriginID,
"origin_url": route.OriginURL,
"origin_host": route.OriginHost,
"upstreams": route.Upstreams,
"enabled": route.Enabled,
"enable_https": route.EnableHTTPS,
"cert_id": route.CertID,
"cert_ids": route.CertIDs,
"domain_cert_ids": route.DomainCertIDs,
"redirect_http": route.RedirectHTTP,
"limit_conn_per_server": route.LimitConnPerServer,
"limit_conn_per_ip": route.LimitConnPerIP,
"limit_rate": route.LimitRate,
"cache_enabled": route.CacheEnabled,
"cache_policy": route.CachePolicy,
"cache_rules": route.CacheRules,
"custom_headers": route.CustomHeaders,
"basic_auth_enabled": route.BasicAuthEnabled,
"basic_auth_username": route.BasicAuthUsername,
"basic_auth_password": route.BasicAuthPassword,
"remark": route.Remark,
"upstream_type": route.UpstreamType,
"tunnel_node_id": route.TunnelNodeID,
"tunnel_target_addr": route.TunnelTargetAddr,
"tunnel_target_protocol": route.TunnelTargetProtocol,
"pages_project_id": route.PagesProjectID,
}).Error
}
func (route *ProxyRoute) Delete() error {
return DB.Delete(route).Error
}
@@ -0,0 +1,60 @@
package model
import (
schemagoose "github.com/rain-kl/openflare/openflare-server/internal/model/goose"
"gorm.io/gorm"
)
func currentGooseTargetVersion() int64 {
return schemagoose.CurrentTargetVersion()
}
func loadGooseDatabaseVersion(db *gorm.DB) (int, bool, error) {
return schemagoose.LoadDatabaseVersion(db)
}
func ensureDatabaseSchemaUpToDate(db *gorm.DB, backend string) error {
return schemagoose.EnsureDatabaseSchemaUpToDate(db, backend, databaseSchemaMigrationContext{})
}
func (databaseSchemaMigrationContext) RegisterSharding(db *gorm.DB, backend string) error {
return registerSharding(db, backend)
}
func (databaseSchemaMigrationContext) AutoMigrateLegacySchemaMetadata(db *gorm.DB) error {
return autoMigrateLegacySchemaMetadata(db)
}
func (databaseSchemaMigrationContext) InitializeFreshDatabaseSchema(db *gorm.DB, backend string) error {
return initializeFreshDatabaseSchema(db, backend)
}
func (databaseSchemaMigrationContext) IsDatabaseEmpty(db *gorm.DB) (bool, error) {
return isDatabaseEmpty(db)
}
func (databaseSchemaMigrationContext) RepairCurrentSchemaState(db *gorm.DB, backend string) error {
if err := dropLegacyNodeColumns(db, backend); err != nil {
return err
}
if err := ensureDefaultGitHubAuthSource(db); err != nil {
return err
}
if err := ensureDefaultWAFRuleGroup(db); err != nil {
return err
}
return nil
}
func (databaseSchemaMigrationContext) SaveLegacyDatabaseSchemaVersion(db *gorm.DB, version int) error {
return saveLegacyDatabaseSchemaVersion(db, version)
}
func (databaseSchemaMigrationContext) UpgradeLegacyDatabaseSchema(db *gorm.DB, backend string, version int) error {
return upgradeLegacyDatabaseSchema(db, backend, version)
}
func (databaseSchemaMigrationContext) ValidateCurrentDatabaseSchema(db *gorm.DB, backend string) error {
return validateCurrentDatabaseSchema(db, backend)
}
+208
View File
@@ -0,0 +1,208 @@
package model
import (
"fmt"
"strconv"
"strings"
"sync"
"github.com/bwmarrin/snowflake"
"gorm.io/gorm"
"gorm.io/sharding"
)
const observabilityShardCount = 10
var (
observabilityIDNode *snowflake.Node
observabilityIDNodeErr error
observabilityIDNodeOnce sync.Once
)
func registerSharding(db *gorm.DB, backend string) error {
if db == nil {
return nil
}
_ = backend
if err := db.Use(sharding.Register(sharding.Config{
ShardingKey: "id",
NumberOfShards: observabilityShardCount,
ShardingAlgorithm: func(value any) (string, error) {
return observabilityShardSuffixForValue(value)
},
ShardingAlgorithmByPrimaryKey: func(id int64) string {
return observabilityShardSuffixForInt64(id)
},
PrimaryKeyGenerator: sharding.PKCustom,
PrimaryKeyGeneratorFn: func(tableIdx int64) int64 {
return 0
},
}, shardedObservabilityTables()...)); err != nil {
return fmt.Errorf("register observability sharding failed: %w", err)
}
return nil
}
func shardedObservabilityTables() []any {
return []any{
&NodeMetricSnapshot{},
&NodeRequestReport{},
&NodeAccessLog{},
&NodeObservationOpenresty{},
&NodeObservationFrps{},
&NodeObservationFrpc{},
}
}
func shardedObservabilityBaseTables() []string {
return []string{
"node_metric_snapshots",
"node_request_reports",
"node_access_logs",
"node_observation_openresties",
"node_observation_frps",
"node_observation_frpcs",
}
}
func isShardedObservabilityTable(tableName string) bool {
switch strings.TrimSpace(tableName) {
case "node_metric_snapshots", "node_request_reports", "node_access_logs", "node_observation_openresties", "node_observation_frps", "node_observation_frpcs":
return true
default:
return false
}
}
func observabilityShardTables(baseTable string) []string {
tables := make([]string, 0, observabilityShardCount)
for _, suffix := range observabilityShardSuffixes() {
tables = append(tables, baseTable+suffix)
}
return tables
}
func observabilityShardSuffixes() []string {
suffixes := make([]string, 0, observabilityShardCount)
for index := 0; index < observabilityShardCount; index++ {
suffixes = append(suffixes, fmt.Sprintf("_%02d", index))
}
return suffixes
}
func observabilityShardSuffixForID(id uint) string {
return fmt.Sprintf("_%02d", uint64(id)%uint64(observabilityShardCount))
}
func observabilityShardSuffixForInt64(id int64) string {
if id < 0 {
id = -id
}
return fmt.Sprintf("_%02d", uint64(id)%uint64(observabilityShardCount))
}
func observabilityShardSuffixForValue(value any) (string, error) {
switch typed := value.(type) {
case int:
return observabilityShardSuffixForInt64(int64(typed)), nil
case int8:
return observabilityShardSuffixForInt64(int64(typed)), nil
case int16:
return observabilityShardSuffixForInt64(int64(typed)), nil
case int32:
return observabilityShardSuffixForInt64(int64(typed)), nil
case int64:
return observabilityShardSuffixForInt64(typed), nil
case uint:
return observabilityShardSuffixForID(typed), nil
case uint8:
return observabilityShardSuffixForID(uint(typed)), nil
case uint16:
return observabilityShardSuffixForID(uint(typed)), nil
case uint32:
return observabilityShardSuffixForID(uint(typed)), nil
case uint64:
return fmt.Sprintf("_%02d", typed%uint64(observabilityShardCount)), nil
case string:
id, err := strconv.ParseUint(strings.TrimSpace(typed), 10, 64)
if err != nil {
return "", fmt.Errorf("invalid sharding id %q", typed)
}
return fmt.Sprintf("_%02d", id%uint64(observabilityShardCount)), nil
default:
return "", fmt.Errorf("unsupported observability sharding value type %T", value)
}
}
func legacyObservabilityShardTableName(tableName string) string {
return tableName + "_legacy_v2_to_v3"
}
func normalizeShardedDB(db *gorm.DB) *gorm.DB {
if db != nil {
return db
}
return DB
}
func nextObservabilityID() (uint, error) {
observabilityIDNodeOnce.Do(func() {
observabilityIDNode, observabilityIDNodeErr = snowflake.NewNode(0)
})
if observabilityIDNodeErr != nil {
return 0, observabilityIDNodeErr
}
id := observabilityIDNode.Generate().Int64()
if id <= 0 {
return 0, fmt.Errorf("generated invalid observability id %d", id)
}
return uint(id), nil
}
func assignObservabilityID(id *uint) error {
if id == nil || *id != 0 {
return nil
}
generated, err := nextObservabilityID()
if err != nil {
return err
}
*id = generated
return nil
}
func queryAcrossShards[T any](baseTable string, query func(tx *gorm.DB) ([]T, error)) ([]T, error) {
return queryAcrossShardsWithDB(DB, baseTable, query)
}
func queryAcrossShardsWithDB[T any](db *gorm.DB, baseTable string, query func(tx *gorm.DB) ([]T, error)) ([]T, error) {
items := make([]T, 0)
db = normalizeShardedDB(db)
for _, table := range observabilityShardTables(baseTable) {
rows, err := query(db.Table(table))
if err != nil {
return nil, err
}
items = append(items, rows...)
}
return items, nil
}
func deleteAcrossShards(db *gorm.DB, baseTable string, model any, apply func(tx *gorm.DB) *gorm.DB) (int64, error) {
db = normalizeShardedDB(db)
var deleted int64
for _, table := range observabilityShardTables(baseTable) {
tx := db.Table(table)
if apply != nil {
tx = apply(tx)
} else {
tx = tx.Session(&gorm.Session{AllowGlobalUpdate: true})
}
result := tx.Delete(model)
if result.Error != nil {
return deleted, result.Error
}
deleted += result.RowsAffected
}
return deleted, nil
}
@@ -0,0 +1,51 @@
package model
import "time"
type TLSCertificate struct {
ID uint `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"uniqueIndex;size:255;not null"`
CertPEM string `json:"-" gorm:"type:text;not null"`
KeyPEM string `json:"-" gorm:"type:text;not null"`
NotBefore time.Time `json:"not_before"`
NotAfter time.Time `json:"not_after"`
Remark string `json:"remark" gorm:"size:255"`
Provider string `json:"provider" gorm:"size:64;default:'upload'"` // upload, acme
AcmeAccountID uint `json:"acme_account_id"`
DnsAccountID uint `json:"dns_account_id"`
KeyAlgorithm string `json:"key_algorithm" gorm:"size:32"`
AutoRenew bool `json:"auto_renew"`
PrimaryDomain string `json:"primary_domain" gorm:"size:255"`
OtherDomains string `json:"other_domains" gorm:"type:text"`
DisableCNAME bool `json:"disable_cname"`
SkipDNS bool `json:"skip_dns"`
DNS1 string `json:"dns1" gorm:"size:128"`
DNS2 string `json:"dns2" gorm:"size:128"`
ApplyStatus string `json:"apply_status" gorm:"size:64;default:'ready'"`
ApplyMessage string `json:"apply_message" gorm:"type:text"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func ListTLSCertificates() (certificates []*TLSCertificate, err error) {
err = DB.Order("id desc").Find(&certificates).Error
return certificates, err
}
func GetTLSCertificateByID(id uint) (*TLSCertificate, error) {
certificate := &TLSCertificate{}
err := DB.First(certificate, id).Error
return certificate, err
}
func (certificate *TLSCertificate) Insert() error {
return DB.Create(certificate).Error
}
func (certificate *TLSCertificate) Update() error {
return DB.Save(certificate).Error
}
func (certificate *TLSCertificate) Delete() error {
return DB.Delete(certificate).Error
}
+193
View File
@@ -0,0 +1,193 @@
package model
import (
"errors"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/utils/security"
)
// User if you add sensitive fields, don't forget to clean them in setupLogin function.
// Otherwise, the sensitive information will be saved on local storage in plain text!
type User struct {
Id int `json:"id"`
Username string `json:"username" gorm:"unique;index" validate:"max=12"`
Password string `json:"password" gorm:"not null;" validate:"min=8,max=20"`
DisplayName string `json:"display_name" gorm:"index" validate:"max=20"`
Role int `json:"role" gorm:"type:int;default:1"` // admin, common
Status int `json:"status" gorm:"type:int;default:1"` // enabled, disabled
Token string `json:"token" gorm:"index"`
Email string `json:"email" gorm:"index" validate:"max=50"`
GitHubId string `json:"github_id" gorm:"column:github_id;index"`
WeChatId string `json:"wechat_id" gorm:"column:wechat_id;index"`
VerificationCode string `json:"verification_code" gorm:"-:all"` // this field is only for Email verification, don't save it to database!
}
func GetMaxUserId() int {
var user User
DB.Last(&user)
return user.Id
}
func GetAllUsers(startIdx int, num int) (users []*User, err error) {
err = DB.Order("id desc").Limit(num).Offset(startIdx).Select([]string{"id", "username", "display_name", "role", "status", "email"}).Find(&users).Error
return users, err
}
func SearchUsers(keyword string) (users []*User, err error) {
err = DB.Select([]string{"id", "username", "display_name", "role", "status", "email"}).Where("id = ? or username LIKE ? or email LIKE ? or display_name LIKE ?", keyword, keyword+"%", keyword+"%", keyword+"%").Find(&users).Error
return users, err
}
func GetUserById(id int, selectAll bool) (*User, error) {
if id == 0 {
return nil, errors.New("id 为空!")
}
user := User{Id: id}
var err error = nil
if selectAll {
err = DB.First(&user, "id = ?", id).Error
} else {
err = DB.Select([]string{"id", "username", "display_name", "role", "status", "email", "wechat_id", "github_id"}).First(&user, "id = ?", id).Error
}
return &user, err
}
func DeleteUserById(id int) (err error) {
if id == 0 {
return errors.New("id 为空!")
}
user := User{Id: id}
return user.Delete()
}
func (user *User) Insert() error {
var err error
if user.Password != "" {
user.Password, err = security.Password2Hash(user.Password)
if err != nil {
return err
}
}
err = DB.Create(user).Error
return err
}
func (user *User) Update(updatePassword bool) error {
var err error
if updatePassword {
user.Password, err = security.Password2Hash(user.Password)
if err != nil {
return err
}
}
err = DB.Model(user).Updates(user).Error
return err
}
func (user *User) Delete() error {
if user.Id == 0 {
return errors.New("id 为空!")
}
err := DB.Delete(user).Error
return err
}
// ValidateAndFill check password & user status
func (user *User) ValidateAndFill() (err error) {
// When querying with struct, GORM will only query with non-zero fields,
// that means if your field’s value is 0, '', false or other zero values,
// it won’t be used to build query conditions
password := user.Password
if user.Username == "" || password == "" {
return errors.New("用户名或密码为空")
}
DB.Where(User{Username: user.Username}).First(user)
okay := security.ValidatePasswordAndHash(password, user.Password)
if !okay || user.Status != common.UserStatusEnabled {
return errors.New("用户名或密码错误,或用户已被封禁")
}
return nil
}
func (user *User) FillUserById() error {
if user.Id == 0 {
return errors.New("id 为空!")
}
DB.Where(User{Id: user.Id}).First(user)
return nil
}
func (user *User) FillUserByEmail() error {
if user.Email == "" {
return errors.New("email 为空!")
}
DB.Where(User{Email: user.Email}).First(user)
return nil
}
func (user *User) FillUserByGitHubId() error {
if user.GitHubId == "" {
return errors.New("GitHub id 为空!")
}
DB.Where(User{GitHubId: user.GitHubId}).First(user)
return nil
}
func (user *User) FillUserByWeChatId() error {
if user.WeChatId == "" {
return errors.New("WeChat id 为空!")
}
DB.Where(User{WeChatId: user.WeChatId}).First(user)
return nil
}
func (user *User) FillUserByUsername() error {
if user.Username == "" {
return errors.New("username 为空!")
}
DB.Where(User{Username: user.Username}).First(user)
return nil
}
// ValidateUserToken looks up a user by their stored JWT token string.
// JWT signature verification is handled by middleware/auth.go; this
// function is used by Logout to find and clear the token from DB.
func ValidateUserToken(token string) (user *User) {
if token == "" {
return nil
}
user = &User{}
if DB.Where("token = ?", token).First(user).RowsAffected == 1 {
return user
}
return nil
}
func IsEmailAlreadyTaken(email string) bool {
return DB.Where("email = ?", email).Find(&User{}).RowsAffected == 1
}
func IsWeChatIdAlreadyTaken(wechatId string) bool {
return DB.Where("wechat_id = ?", wechatId).Find(&User{}).RowsAffected == 1
}
func IsGitHubIdAlreadyTaken(githubId string) bool {
return DB.Where("github_id = ?", githubId).Find(&User{}).RowsAffected == 1
}
func IsUsernameAlreadyTaken(username string) bool {
return DB.Where("username = ?", username).Find(&User{}).RowsAffected == 1
}
func ResetUserPasswordByEmail(email string, password string) error {
if email == "" || password == "" {
return errors.New("邮箱地址或密码为空!")
}
hashedPassword, err := security.Password2Hash(password)
if err != nil {
return err
}
err = DB.Model(&User{}).Where("email = ?", email).Update("password", hashedPassword).Error
return err
}
+168
View File
@@ -0,0 +1,168 @@
package model
import "time"
type WAFRuleGroup struct {
ID uint `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"size:255;not null"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
IsGlobal bool `json:"is_global" gorm:"not null;default:false;index"`
BlockStatusCode int `json:"block_status_code" gorm:"not null;default:418"`
BlockResponseBody string `json:"block_response_body" gorm:"type:text;not null;default:''"`
IPWhitelist string `json:"ip_whitelist" gorm:"type:text;not null;default:'[]'"`
IPBlacklist string `json:"ip_blacklist" gorm:"type:text;not null;default:'[]'"`
IPWhitelistGroups string `json:"ip_whitelist_group_ids" gorm:"type:text;not null;default:'[]'"`
IPBlacklistGroups string `json:"ip_blacklist_group_ids" gorm:"type:text;not null;default:'[]'"`
CountryWhitelist string `json:"country_whitelist" gorm:"type:text;not null;default:'[]'"`
CountryBlacklist string `json:"country_blacklist" gorm:"type:text;not null;default:'[]'"`
RegionWhitelist string `json:"region_whitelist" gorm:"type:text;not null;default:'[]'"`
RegionBlacklist string `json:"region_blacklist" gorm:"type:text;not null;default:'[]'"`
PoWEnabled bool `json:"pow_enabled" gorm:"column:pow_enabled;not null;default:false"`
PoWConfig string `json:"pow_config" gorm:"column:pow_config;type:text;not null;default:'{}'"`
Remark string `json:"remark" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type WAFIPGroup struct {
ID uint `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"size:255;not null"`
Type string `json:"type" gorm:"size:32;not null;index"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
IPList string `json:"ip_list" gorm:"type:text;not null;default:'[]'"`
AutoConfig string `json:"auto_config" gorm:"type:text;not null;default:'{}'"`
ExtIPs string `json:"ext_ips" gorm:"type:text;not null;default:'[]'"`
SubscriptionURL string `json:"subscription_url" gorm:"size:2048;not null;default:''"`
SubscriptionFormat string `json:"subscription_format" gorm:"size:32;not null;default:'text'"`
SubscriptionMappingRule string `json:"subscription_mapping_rule" gorm:"size:255;not null;default:''"`
SyncIntervalMinutes int `json:"sync_interval_minutes" gorm:"not null;default:1440"`
LastSyncedAt *time.Time `json:"last_synced_at"`
NextSyncAt *time.Time `json:"next_sync_at" gorm:"index"`
LastSyncStatus string `json:"last_sync_status" gorm:"size:32;not null;default:''"`
LastSyncMessage string `json:"last_sync_message" gorm:"type:text;not null;default:''"`
Remark string `json:"remark" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type WAFRuleGroupBinding struct {
ID uint `json:"id" gorm:"primaryKey"`
RuleGroupID uint `json:"rule_group_id" gorm:"not null;uniqueIndex:idx_waf_group_route"`
ProxyRouteID uint `json:"proxy_route_id" gorm:"not null;uniqueIndex:idx_waf_group_route;index"`
CreatedAt time.Time `json:"created_at"`
}
func ListWAFRuleGroups() ([]*WAFRuleGroup, error) {
var groups []*WAFRuleGroup
err := DB.Order("is_global desc").Order("id asc").Find(&groups).Error
return groups, err
}
func GetWAFRuleGroupByID(id uint) (*WAFRuleGroup, error) {
group := &WAFRuleGroup{}
err := DB.First(group, id).Error
return group, err
}
func GetGlobalWAFRuleGroup() (*WAFRuleGroup, error) {
group := &WAFRuleGroup{}
err := DB.Where("is_global = ?", true).Order("id asc").First(group).Error
return group, err
}
func (group *WAFRuleGroup) Insert() error {
return DB.Create(group).Error
}
func (group *WAFRuleGroup) Update() error {
return DB.Model(&WAFRuleGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
"name": group.Name,
"enabled": group.Enabled,
"is_global": group.IsGlobal,
"block_status_code": group.BlockStatusCode,
"block_response_body": group.BlockResponseBody,
"ip_whitelist": group.IPWhitelist,
"ip_blacklist": group.IPBlacklist,
"ip_whitelist_groups": group.IPWhitelistGroups,
"ip_blacklist_groups": group.IPBlacklistGroups,
"country_whitelist": group.CountryWhitelist,
"country_blacklist": group.CountryBlacklist,
"region_whitelist": group.RegionWhitelist,
"region_blacklist": group.RegionBlacklist,
"pow_enabled": group.PoWEnabled,
"pow_config": group.PoWConfig,
"remark": group.Remark,
}).Error
}
func (group *WAFRuleGroup) Delete() error {
return DB.Delete(group).Error
}
func ListWAFIPGroups() ([]*WAFIPGroup, error) {
var groups []*WAFIPGroup
err := DB.Order("type asc").Order("id asc").Find(&groups).Error
return groups, err
}
func GetWAFIPGroupByID(id uint) (*WAFIPGroup, error) {
group := &WAFIPGroup{}
err := DB.First(group, id).Error
return group, err
}
func ListWAFIPGroupsByIDs(ids []uint) ([]*WAFIPGroup, error) {
if len(ids) == 0 {
return []*WAFIPGroup{}, nil
}
var groups []*WAFIPGroup
err := DB.Where("id IN ?", ids).Order("id asc").Find(&groups).Error
return groups, err
}
func ListDueWAFIPGroups(now time.Time) ([]*WAFIPGroup, error) {
var groups []*WAFIPGroup
err := DB.Where("enabled = ? AND (type = ? OR (type = ? AND subscription_url <> '')) AND (next_sync_at IS NULL OR next_sync_at <= ?)", true, "automatic", "subscription", now).
Order("id asc").
Find(&groups).Error
return groups, err
}
func (group *WAFIPGroup) Insert() error {
return DB.Create(group).Error
}
func (group *WAFIPGroup) Update() error {
return DB.Model(&WAFIPGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
"name": group.Name,
"type": group.Type,
"enabled": group.Enabled,
"ip_list": group.IPList,
"auto_config": group.AutoConfig,
"ext_ips": group.ExtIPs,
"subscription_url": group.SubscriptionURL,
"subscription_format": group.SubscriptionFormat,
"subscription_mapping_rule": group.SubscriptionMappingRule,
"sync_interval_minutes": group.SyncIntervalMinutes,
"next_sync_at": group.NextSyncAt,
"last_sync_status": group.LastSyncStatus,
"last_sync_message": group.LastSyncMessage,
"remark": group.Remark,
}).Error
}
func (group *WAFIPGroup) UpdateSyncResult() error {
return DB.Model(&WAFIPGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
"ip_list": group.IPList,
"ext_ips": group.ExtIPs,
"last_synced_at": group.LastSyncedAt,
"next_sync_at": group.NextSyncAt,
"last_sync_status": group.LastSyncStatus,
"last_sync_message": group.LastSyncMessage,
"subscription_format": group.SubscriptionFormat,
}).Error
}
func (group *WAFIPGroup) Delete() error {
return DB.Delete(group).Error
}
@@ -0,0 +1,271 @@
package router
import (
"github.com/rain-kl/openflare/openflare-server/internal/controller"
"github.com/rain-kl/openflare/openflare-server/internal/middleware"
"github.com/gin-gonic/gin"
)
func SetApiRouter(router *gin.Engine) {
apiRouter := router.Group("/api")
apiRouter.Use(middleware.GlobalAPIRateLimit())
{
apiRouter.GET("/status", controller.GetStatus)
apiRouter.GET("/notice", controller.GetNotice)
apiRouter.GET("/about", controller.GetAbout)
apiRouter.GET("/verification", middleware.CriticalRateLimit(), controller.SendEmailVerification)
apiRouter.GET("/reset_password", middleware.CriticalRateLimit(), controller.SendPasswordResetEmail)
apiRouter.POST("/user/reset", middleware.CriticalRateLimit(), controller.ResetPassword)
apiRouter.GET("/oauth/github", middleware.CriticalRateLimit(), controller.GitHubOAuth)
apiRouter.GET("/oauth/wechat", middleware.CriticalRateLimit(), controller.WeChatAuth)
apiRouter.GET("/oauth/wechat/bind", middleware.CriticalRateLimit(), middleware.UserAuth(), controller.WeChatBind)
apiRouter.GET("/oauth/email/bind", middleware.CriticalRateLimit(), middleware.UserAuth(), controller.EmailBind)
apiRouter.GET("/oauth/:source/authorize", middleware.CriticalRateLimit(), controller.OAuthAuthorize)
apiRouter.GET("/oauth/:source/callback", middleware.CriticalRateLimit(), controller.OAuthCallback)
apiRouter.POST("/oauth/link-existing", middleware.CriticalRateLimit(), controller.LinkExistingOAuthAccount)
externalAccountRoute := apiRouter.Group("/oauth/external-accounts")
externalAccountRoute.Use(middleware.UserAuth(), middleware.NoTokenAuth())
{
externalAccountRoute.GET("/", controller.ListExternalAccounts)
externalAccountRoute.POST("/:id/delete", controller.DeleteExternalAccount)
}
capRoute := apiRouter.Group("/cap")
{
capRoute.POST("/:scope/challenge", middleware.CriticalRateLimit(), controller.GetCapChallenge)
capRoute.POST("/:scope/redeem", middleware.CriticalRateLimit(), controller.RedeemCapChallenge)
}
userRoute := apiRouter.Group("/user")
{
userRoute.POST("/register", middleware.CriticalRateLimit(), controller.Register)
userRoute.POST("/login", middleware.CriticalRateLimit(), middleware.CapAuth("login"), controller.Login)
userRoute.GET("/logout", controller.Logout)
selfRoute := userRoute.Group("/")
selfRoute.Use(middleware.UserAuth(), middleware.NoTokenAuth())
{
selfRoute.GET("/self", controller.GetSelf)
selfRoute.POST("/self/update", controller.UpdateSelf)
selfRoute.POST("/self/delete", controller.DeleteSelf)
selfRoute.GET("/token", controller.GenerateToken)
}
adminRoute := userRoute.Group("/")
adminRoute.Use(middleware.AdminAuth(), middleware.NoTokenAuth())
{
adminRoute.GET("/", controller.GetAllUsers)
adminRoute.GET("/search", controller.SearchUsers)
adminRoute.GET("/:id", controller.GetUser)
adminRoute.POST("/", controller.CreateUser)
adminRoute.POST("/manage", controller.ManageUser)
adminRoute.POST("/update", controller.UpdateUser)
adminRoute.POST("/:id/delete", controller.DeleteUser)
}
}
optionRoute := apiRouter.Group("/option")
optionRoute.Use(middleware.RootAuth(), middleware.NoTokenAuth())
{
optionRoute.GET("/", controller.GetOptions)
optionRoute.POST("/update", controller.UpdateOption)
optionRoute.POST("/update-batch", controller.UpdateOptionsBatch)
optionRoute.POST("/geoip/lookup", controller.LookupGeoIP)
optionRoute.POST("/database/cleanup", controller.CleanupDatabaseObservability)
}
uptimekumaRoute := apiRouter.Group("/uptimekuma")
uptimekumaRoute.Use(middleware.RootAuth(), middleware.NoTokenAuth())
{
uptimekumaRoute.POST("/sync", controller.SyncUptimeKuma)
}
authSourceRoute := apiRouter.Group("/auth-sources")
authSourceRoute.Use(middleware.RootAuth(), middleware.NoTokenAuth())
{
authSourceRoute.GET("/", controller.ListAuthSources)
authSourceRoute.POST("/", controller.CreateAuthSource)
authSourceRoute.POST("/:id/update", controller.UpdateAuthSource)
authSourceRoute.POST("/:id/delete", controller.DeleteAuthSource)
authSourceRoute.POST("/:id/toggle", controller.ToggleAuthSource)
}
updateRoute := apiRouter.Group("/update")
updateRoute.Use(middleware.RootAuth(), middleware.NoTokenAuth())
{
updateRoute.GET("/latest-release", controller.GetLatestRelease)
updateRoute.GET("/logs/ws", controller.StreamServerUpgradeLogs)
updateRoute.POST("/manual-upload", controller.UploadManualServerBinary)
updateRoute.POST("/manual-upgrade", controller.ConfirmManualServerUpgrade)
updateRoute.POST("/upgrade", controller.UpgradeServer)
}
proxyRoute := apiRouter.Group("/proxy-routes")
proxyRoute.Use(middleware.AdminAuth())
{
proxyRoute.GET("/", controller.GetProxyRoutes)
proxyRoute.GET("/:id", controller.GetProxyRoute)
proxyRoute.POST("/", controller.CreateProxyRoute)
proxyRoute.POST("/:id/update", controller.UpdateProxyRoute)
proxyRoute.POST("/:id/delete", controller.DeleteProxyRoute)
}
wafRoute := apiRouter.Group("/waf")
wafRoute.Use(middleware.AdminAuth())
{
wafRoute.GET("/ip-groups", controller.ListWAFIPGroups)
wafRoute.GET("/ip-groups/:id", controller.GetWAFIPGroup)
wafRoute.POST("/ip-groups", controller.CreateWAFIPGroup)
wafRoute.POST("/ip-groups/test", controller.TestWAFIPGroupAutoConfig)
wafRoute.POST("/ip-groups/:id/update", controller.UpdateWAFIPGroup)
wafRoute.POST("/ip-groups/:id/delete", controller.DeleteWAFIPGroup)
wafRoute.POST("/ip-groups/:id/sync", controller.SyncWAFIPGroup)
wafRoute.GET("/rule-groups", controller.ListWAFRuleGroups)
wafRoute.GET("/rule-groups/:id", controller.GetWAFRuleGroup)
wafRoute.POST("/rule-groups", controller.CreateWAFRuleGroup)
wafRoute.POST("/rule-groups/:id/update", controller.UpdateWAFRuleGroup)
wafRoute.POST("/rule-groups/:id/delete", controller.DeleteWAFRuleGroup)
wafRoute.POST("/rule-groups/:id/sites", controller.ReplaceWAFRuleGroupSites)
wafRoute.GET("/sites/:route_id/rule-groups", controller.GetWAFSiteRuleGroups)
wafRoute.POST("/sites/:route_id/rule-groups", controller.ReplaceWAFSiteRuleGroups)
}
originRoute := apiRouter.Group("/origins")
originRoute.Use(middleware.AdminAuth())
{
originRoute.GET("/", controller.GetOrigins)
originRoute.GET("/:id", controller.GetOrigin)
originRoute.POST("/", controller.CreateOrigin)
originRoute.POST("/:id/update", controller.UpdateOrigin)
originRoute.POST("/:id/delete", controller.DeleteOrigin)
}
pagesRoute := apiRouter.Group("/pages")
pagesRoute.Use(middleware.AdminAuth())
{
pagesRoute.GET("/", controller.ListPagesProjects)
pagesRoute.GET("/:id", controller.GetPagesProject)
pagesRoute.POST("/", controller.CreatePagesProject)
pagesRoute.POST("/:id/update", controller.UpdatePagesProject)
pagesRoute.POST("/:id/delete", controller.DeletePagesProject)
pagesRoute.GET("/:id/deployments", controller.ListPagesDeployments)
pagesRoute.POST("/:id/deployments/upload", controller.UploadPagesDeployment)
pagesRoute.POST("/:id/deployments/:deployment_id/activate", controller.ActivatePagesDeployment)
pagesRoute.POST("/:id/deployments/:deployment_id/delete", controller.DeletePagesDeployment)
pagesRoute.GET("/deployments/:deployment_id/files", controller.ListPagesDeploymentFiles)
}
managedDomainRoute := apiRouter.Group("/managed-domains")
managedDomainRoute.Use(middleware.AdminAuth())
{
managedDomainRoute.GET("/", controller.GetManagedDomains)
managedDomainRoute.GET("/match", controller.MatchManagedDomainCertificate)
managedDomainRoute.POST("/", controller.CreateManagedDomain)
managedDomainRoute.POST("/:id/update", controller.UpdateManagedDomain)
managedDomainRoute.POST("/:id/delete", controller.DeleteManagedDomain)
}
tlsCertificateRoute := apiRouter.Group("/tls-certificates")
tlsCertificateRoute.Use(middleware.AdminAuth())
{
tlsCertificateRoute.GET("/", controller.GetTLSCertificates)
tlsCertificateRoute.GET("/:id", controller.GetTLSCertificate)
tlsCertificateRoute.GET("/:id/content", controller.GetTLSCertificateContent)
tlsCertificateRoute.POST("/", controller.CreateTLSCertificate)
tlsCertificateRoute.POST("/:id/update", controller.UpdateTLSCertificate)
tlsCertificateRoute.POST("/:id/update-acme", controller.UpdateAcmeCertificate)
tlsCertificateRoute.POST("/:id/convert-acme", controller.ConvertTLSCertificateToAcme)
tlsCertificateRoute.POST("/import-file", controller.ImportTLSCertificateFile)
tlsCertificateRoute.POST("/:id/delete", controller.DeleteTLSCertificate)
tlsCertificateRoute.POST("/apply", controller.ApplyTLSCertificate)
tlsCertificateRoute.POST("/:id/renew", controller.RenewTLSCertificate)
}
acmeAccountRoute := apiRouter.Group("/acme-accounts")
acmeAccountRoute.Use(middleware.AdminAuth())
{
acmeAccountRoute.GET("/default", controller.GetDefaultAcmeAccount)
}
dnsAccountRoute := apiRouter.Group("/dns-accounts")
dnsAccountRoute.Use(middleware.AdminAuth())
{
dnsAccountRoute.GET("/", controller.GetDnsAccounts)
dnsAccountRoute.POST("/", controller.CreateDnsAccount)
dnsAccountRoute.POST("/:id/update", controller.UpdateDnsAccount)
dnsAccountRoute.POST("/:id/delete", controller.DeleteDnsAccount)
}
configVersionRoute := apiRouter.Group("/config-versions")
configVersionRoute.Use(middleware.AdminAuth())
{
configVersionRoute.GET("/", controller.GetConfigVersions)
configVersionRoute.GET("/active", controller.GetActiveConfigVersion)
configVersionRoute.GET("/preview", controller.PreviewConfigVersion)
configVersionRoute.GET("/diff", controller.DiffConfigVersion)
configVersionRoute.GET("/:id", controller.GetConfigVersion)
configVersionRoute.POST("/publish", controller.PublishConfigVersion)
configVersionRoute.POST("/:id/activate", controller.ActivateConfigVersion)
configVersionRoute.POST("/cleanup", controller.CleanupConfigVersions)
}
dashboardRoute := apiRouter.Group("/dashboard")
dashboardRoute.Use(middleware.AdminAuth())
{
dashboardRoute.GET("/overview", controller.GetDashboardOverview)
}
nodeRoute := apiRouter.Group("/nodes")
nodeRoute.Use(middleware.AdminAuth())
{
nodeRoute.GET("/bootstrap-token", controller.GetNodeBootstrapToken)
nodeRoute.POST("/bootstrap-token/rotate", controller.RotateNodeBootstrapToken)
nodeRoute.GET("/", controller.GetNodes)
nodeRoute.POST("/", controller.CreateNode)
nodeRoute.GET("/:id/agent-release", controller.GetNodeAgentRelease)
nodeRoute.POST("/:id/update", controller.UpdateNode)
nodeRoute.POST("/:id/delete", controller.DeleteNode)
nodeRoute.POST("/:id/agent-update", controller.RequestNodeAgentUpdate)
nodeRoute.POST("/:id/openresty-restart", controller.RequestNodeOpenrestyRestart)
nodeRoute.POST("/:id/force-sync", controller.RequestNodeForceSync)
nodeRoute.GET("/:id/observability", controller.GetNodeObservability)
nodeRoute.POST("/:id/observability/cleanup", controller.CleanupNodeHealthEvents)
}
applyLogRoute := apiRouter.Group("/apply-logs")
applyLogRoute.Use(middleware.AdminAuth())
{
applyLogRoute.GET("/", controller.GetApplyLogs)
applyLogRoute.POST("/cleanup", controller.CleanupApplyLogs)
}
accessLogRoute := apiRouter.Group("/access-logs")
accessLogRoute.Use(middleware.AdminAuth())
{
accessLogRoute.GET("/", controller.GetAccessLogs)
accessLogRoute.GET("/folds", controller.GetFoldedAccessLogs)
accessLogRoute.GET("/folds/ip-summary", controller.GetFoldedAccessLogIPs)
accessLogRoute.GET("/ip-summary", controller.GetAccessLogIPSummaries)
accessLogRoute.GET("/ip-summary/trend", controller.GetAccessLogIPTrend)
accessLogRoute.POST("/cleanup", controller.CleanupAccessLogs)
}
agentRoute := apiRouter.Group("/agent")
{
discoveryRoute := agentRoute.Group("/")
discoveryRoute.Use(middleware.AgentRegisterAuth())
{
discoveryRoute.POST("/nodes/register", controller.AgentRegister)
}
authorizedRoute := agentRoute.Group("/")
authorizedRoute.Use(middleware.AgentAuth())
{
authorizedRoute.GET("/ws", controller.AgentWebSocket)
authorizedRoute.POST("/nodes/heartbeat", controller.AgentHeartbeat)
authorizedRoute.GET("/config-versions/active", controller.AgentGetActiveConfig)
authorizedRoute.GET("/pages/deployments/:deployment_id/package", controller.AgentDownloadPagesDeploymentPackage)
authorizedRoute.POST("/waf/ip-groups/sync", controller.AgentSyncWAFIPGroups)
authorizedRoute.POST("/apply-logs", controller.AgentReportApplyLog)
}
}
relayRoute := apiRouter.Group("/relay")
relayRoute.Use(middleware.RelayAuth())
{
relayRoute.POST("/heartbeat", controller.RelayHeartbeat)
relayRoute.GET("/ws", controller.RelayWebSocket)
}
flaredRoute := apiRouter.Group("/flared")
flaredRoute.Use(middleware.TunnelAuth())
{
flaredRoute.POST("/heartbeat", controller.FlaredHeartbeat)
flaredRoute.GET("/config/active", controller.FlaredGetActiveConfig)
flaredRoute.POST("/apply-log", controller.FlaredReportApplyLog)
flaredRoute.GET("/ws", controller.FlaredWebSocket)
}
}
}
@@ -0,0 +1,193 @@
package router_test
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/rain-kl/openflare/openflare-server/internal/router"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
)
func TestPhaseFlaredRoutesUnauthorized(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
heartbeatReq := httptest.NewRequest(http.MethodPost, "/api/flared/heartbeat", bytes.NewReader([]byte(`{}`)))
heartbeatReq.Header.Set("Content-Type", "application/json")
heartbeatRec := httptest.NewRecorder()
engine.ServeHTTP(heartbeatRec, heartbeatReq)
if heartbeatRec.Code != http.StatusUnauthorized {
t.Fatalf("expected unauthorized status for missing token, got %d body=%s", heartbeatRec.Code, heartbeatRec.Body.String())
}
activeReq := httptest.NewRequest(http.MethodGet, "/api/flared/config/active", nil)
activeRec := httptest.NewRecorder()
engine.ServeHTTP(activeRec, activeReq)
if activeRec.Code != http.StatusUnauthorized {
t.Fatalf("expected unauthorized status for missing token on active config, got %d", activeRec.Code)
}
applyReq := httptest.NewRequest(http.MethodPost, "/api/flared/apply-log", bytes.NewReader([]byte(`{}`)))
applyReq.Header.Set("Content-Type", "application/json")
applyRec := httptest.NewRecorder()
engine.ServeHTTP(applyRec, applyReq)
if applyRec.Code != http.StatusUnauthorized {
t.Fatalf("expected unauthorized status for missing token on apply log, got %d", applyRec.Code)
}
}
func TestPhaseFlaredRoutesRejectWrongNodeType(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
adminToken := prepareRootToken(t)
createNodeResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/nodes/", map[string]any{
"name": "edge-for-flared-test",
"ip": "10.0.0.20",
})
var createdNode service.NodeView
decodeResponseData(t, createNodeResp, &createdNode)
heartbeatReq := httptest.NewRequest(http.MethodPost, "/api/flared/heartbeat", bytes.NewReader([]byte(`{}`)))
heartbeatReq.Header.Set("Content-Type", "application/json")
heartbeatReq.Header.Set("X-Tunnel-Token", createdNode.AccessToken)
heartbeatRec := httptest.NewRecorder()
engine.ServeHTTP(heartbeatRec, heartbeatReq)
if heartbeatRec.Code != http.StatusForbidden {
t.Fatalf("expected forbidden status for edge_node token, got %d body=%s", heartbeatRec.Code, heartbeatRec.Body.String())
}
}
func TestPhaseFlaredLifecycle(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
adminToken := prepareRootToken(t)
// Create an enabled proxy route that will be served to the flared client
// through the tunnel upstream flow.
createRouteAndPublishVersion(t, engine, adminToken)
// Seed a tunnel_client node directly so we can use its access token as the
// tunnel_token when calling the flared endpoints.
tunnelNode := &model.Node{
NodeID: "tun-flared-1",
Name: "office-flared-1",
IP: "192.168.10.20",
AccessToken: "tunnel-token-phase",
Status: service.NodeStatusPending,
NodeType: "tunnel_client",
Version: "",
}
if err := tunnelNode.Insert(); err != nil {
t.Fatalf("failed to seed tunnel client node: %v", err)
}
heartbeatResp := performFlaredJSONRequest(t, engine, tunnelNode.AccessToken, http.MethodPost, "/api/flared/heartbeat", map[string]any{
"client_version": "v0.2.0",
"frp_version": "0.61.0",
"tunnel_status": "running",
"current_version": "",
})
if !heartbeatResp.Success {
t.Fatalf("flared heartbeat failed: %s", heartbeatResp.Message)
}
var heartbeatData service.FlaredHeartbeatResponse
if err := json.Unmarshal(heartbeatResp.Data, &heartbeatData); err != nil {
t.Fatalf("failed to decode flared heartbeat response: %v", err)
}
if heartbeatData.ActiveConfig == nil {
t.Fatal("expected heartbeat to return active config summary")
}
if heartbeatData.TunnelSettings == nil {
t.Fatal("expected heartbeat to return tunnel_settings")
}
// Re-fetch node and assert status flipped to online.
updated, err := model.GetNodeByNodeID(tunnelNode.NodeID)
if err != nil {
t.Fatalf("failed to reload flared node: %v", err)
}
if updated.Status != service.NodeStatusOnline {
t.Fatalf("expected flared node status to be online, got %q", updated.Status)
}
if updated.Version != "v0.2.0" {
t.Fatalf("expected flared client_version to be stored, got %q", updated.Version)
}
activeResp := performFlaredJSONRequest(t, engine, tunnelNode.AccessToken, http.MethodGet, "/api/flared/config/active", nil)
if !activeResp.Success {
t.Fatalf("flared get active config failed: %s", activeResp.Message)
}
var activeConfig service.FlaredTunnelConfigResponse
if err := json.Unmarshal(activeResp.Data, &activeConfig); err != nil {
t.Fatalf("failed to decode flared active config: %v", err)
}
if activeConfig.Version == "" || activeConfig.Checksum == "" {
t.Fatalf("expected flared active config to return version summary, got %+v", activeConfig)
}
applyResp := performFlaredJSONRequest(t, engine, tunnelNode.AccessToken, http.MethodPost, "/api/flared/apply-log", map[string]any{
"version": activeConfig.Version,
"result": service.ApplyResultOK,
"message": "apply ok",
"checksum": activeConfig.Checksum,
})
if !applyResp.Success {
t.Fatalf("flared apply log failed: %s", applyResp.Message)
}
}
func performFlaredJSONRequest(t *testing.T, engine http.Handler, token string, method string, path string, body any) apiResponse {
t.Helper()
var payload []byte
if body != nil {
var err error
payload, err = json.Marshal(body)
if err != nil {
t.Fatalf("failed to marshal request body: %v", err)
}
}
req := httptest.NewRequest(method, path, bytes.NewReader(payload))
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
req.Header.Set("X-Tunnel-Token", token)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected status %d for %s %s: %s", recorder.Code, method, path, recorder.Body.String())
}
var resp apiResponse
if err := json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if !resp.Success {
t.Fatalf("request %s %s failed: %s", method, path, resp.Message)
}
return resp
}
@@ -0,0 +1,649 @@
package router_test
import (
"bytes"
"crypto/rand"
"crypto/rsa"
"crypto/x509"
"crypto/x509/pkix"
"encoding/json"
"encoding/pem"
"errors"
"math/big"
"mime/multipart"
"net/http"
"net/http/httptest"
"path/filepath"
"strconv"
"strings"
"testing"
"time"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/middleware"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/rain-kl/openflare/openflare-server/internal/router"
"github.com/rain-kl/openflare/openflare-server/internal/service"
)
type apiResponse struct {
Success bool `json:"success"`
Message string `json:"message"`
Data json.RawMessage `json:"data"`
}
func TestPhase1PublishLifecycle(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
token := prepareRootToken(t)
createBody := map[string]any{
"domain": "app.example.com",
"origin_url": "https://10.0.0.11:8443",
"upstreams": []string{"https://10.0.0.12:8443"},
"origin_host": "origin-a.internal",
"enabled": true,
"cache_enabled": true,
"cache_policy": "path_prefix",
"cache_rules": []string{"/assets", "/static"},
"remark": "primary route",
}
resp := performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/", createBody)
var createdRoute service.ProxyRouteView
decodeResponseData(t, resp, &createdRoute)
if createdRoute.Domain != "app.example.com" {
t.Fatalf("unexpected created route domain: %s", createdRoute.Domain)
}
if createdRoute.OriginHost != "origin-a.internal" {
t.Fatalf("unexpected created route origin host: %s", createdRoute.OriginHost)
}
if !createdRoute.CacheEnabled || createdRoute.CachePolicy != "path_prefix" {
t.Fatalf("expected route cache settings to persist, got %+v", createdRoute)
}
if !strings.Contains(createdRoute.Upstreams, "10.0.0.12:8443") {
t.Fatalf("expected route upstream list to persist, got %s", createdRoute.Upstreams)
}
if !strings.Contains(createdRoute.CacheRules, "/assets") {
t.Fatalf("expected route cache rules to persist, got %s", createdRoute.CacheRules)
}
resp = performJSONRequest(t, engine, token, http.MethodGet, "/api/proxy-routes/", nil)
var routes []service.ProxyRouteView
decodeResponseData(t, resp, &routes)
if len(routes) != 1 {
t.Fatalf("expected 1 route, got %d", len(routes))
}
resp = performJSONRequest(t, engine, token, http.MethodPost, "/api/config-versions/publish", nil)
var version1 model.ConfigVersion
decodeResponseData(t, resp, &version1)
if !version1.IsActive {
t.Fatal("expected published version to be active")
}
if version1.SnapshotJSON == "" || version1.RenderedConfig == "" || version1.Checksum == "" {
t.Fatal("expected published version to contain snapshot, rendered config and checksum")
}
if version1.MainConfig == "" {
t.Fatal("expected published version to contain main config")
}
repeatPublishReq := httptest.NewRequest(http.MethodPost, "/api/config-versions/publish", nil)
repeatPublishReq.Header.Set("OpenFlare-Token", token)
repeatPublishRecorder := httptest.NewRecorder()
engine.ServeHTTP(repeatPublishRecorder, repeatPublishReq)
if repeatPublishRecorder.Code != http.StatusOK {
t.Fatalf("unexpected status %d for repeated publish: %s", repeatPublishRecorder.Code, repeatPublishRecorder.Body.String())
}
var repeatPublishResp apiResponse
if err := json.Unmarshal(repeatPublishRecorder.Body.Bytes(), &repeatPublishResp); err != nil {
t.Fatalf("failed to unmarshal repeated publish response: %v", err)
}
if repeatPublishResp.Success {
t.Fatal("expected repeated publish without route changes to be rejected")
}
if !strings.Contains(repeatPublishResp.Message, "当前规则没有变更") {
t.Fatalf("unexpected repeated publish message: %s", repeatPublishResp.Message)
}
initialSnapshot := version1.SnapshotJSON
initialMainConfig := version1.MainConfig
initialRendered := version1.RenderedConfig
updateBody := map[string]any{
"domain": "app.example.com",
"origin_url": "https://10.0.0.21:8443",
"upstreams": []string{"https://10.0.0.22:8443"},
"origin_host": "origin-b.internal",
"enabled": true,
"cache_enabled": true,
"cache_policy": "path_exact",
"cache_rules": []string{"/robots.txt"},
"remark": "updated route",
}
routePath := "/api/proxy-routes/" + toString(createdRoute.ID)
resp = performJSONRequest(t, engine, token, http.MethodPost, routePath+"/update", updateBody)
decodeResponseData(t, resp, &createdRoute)
if createdRoute.OriginURL != "https://10.0.0.21:8443" {
t.Fatalf("unexpected updated route origin: %s", createdRoute.OriginURL)
}
if createdRoute.OriginHost != "origin-b.internal" {
t.Fatalf("unexpected updated route origin host: %s", createdRoute.OriginHost)
}
if createdRoute.CachePolicy != "path_exact" || !strings.Contains(createdRoute.CacheRules, "/robots.txt") {
t.Fatalf("expected updated route cache rules to persist, got %+v", createdRoute)
}
if !strings.Contains(createdRoute.Upstreams, "10.0.0.22:8443") {
t.Fatalf("expected updated route upstream list to persist, got %s", createdRoute.Upstreams)
}
resp = performJSONRequest(t, engine, token, http.MethodPost, "/api/config-versions/publish", nil)
var version2 model.ConfigVersion
decodeResponseData(t, resp, &version2)
if version2.ID == version1.ID {
t.Fatal("expected a new version record")
}
resp = performJSONRequest(t, engine, token, http.MethodGet, "/api/config-versions/", nil)
var versions []map[string]any
decodeResponseData(t, resp, &versions)
if len(versions) != 2 {
t.Fatalf("expected 2 versions, got %d", len(versions))
}
if _, ok := versions[0]["snapshot_json"]; ok {
t.Fatal("expected config version list to omit snapshot_json")
}
if _, ok := versions[0]["main_config"]; ok {
t.Fatal("expected config version list to omit main_config")
}
if _, ok := versions[0]["rendered_config"]; ok {
t.Fatal("expected config version list to omit rendered_config")
}
if _, ok := versions[0]["support_files_json"]; ok {
t.Fatal("expected config version list to omit support_files_json")
}
detailResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/config-versions/"+toString(version2.ID), nil)
var versionDetail model.ConfigVersion
decodeResponseData(t, detailResp, &versionDetail)
if versionDetail.ID != version2.ID {
t.Fatalf("expected config version detail %d, got %d", version2.ID, versionDetail.ID)
}
if versionDetail.SnapshotJSON == "" || versionDetail.MainConfig == "" || versionDetail.RenderedConfig == "" {
t.Fatal("expected config version detail endpoint to include full payload")
}
activeResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/config-versions/active", nil)
var activeVersion model.ConfigVersion
decodeResponseData(t, activeResp, &activeVersion)
if activeVersion.ID != version2.ID {
t.Fatalf("expected version %d active, got %d", version2.ID, activeVersion.ID)
}
activatePath := "/api/config-versions/" + toString(version1.ID) + "/activate"
resp = performJSONRequest(t, engine, token, http.MethodPost, activatePath, nil)
decodeResponseData(t, resp, &activeVersion)
if activeVersion.ID != version1.ID || !activeVersion.IsActive {
t.Fatal("expected version1 to become active after rollback activation")
}
var storedVersion1 model.ConfigVersion
if err := model.DB.First(&storedVersion1, version1.ID).Error; err != nil {
t.Fatalf("failed to query version1: %v", err)
}
if storedVersion1.SnapshotJSON != initialSnapshot {
t.Fatal("expected version1 snapshot to remain immutable")
}
if storedVersion1.MainConfig != initialMainConfig {
t.Fatal("expected version1 main config to remain immutable")
}
if storedVersion1.RenderedConfig != initialRendered {
t.Fatal("expected version1 rendered config to remain immutable")
}
deletePath := "/api/proxy-routes/" + toString(createdRoute.ID)
resp = performJSONRequest(t, engine, token, http.MethodPost, deletePath+"/delete", nil)
if !resp.Success {
t.Fatalf("expected delete route success, got: %s", resp.Message)
}
}
func TestPhase1HTTPSAndCertificateImportLifecycle(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
token := prepareRootToken(t)
certPEM, keyPEM := generateCertificatePairForRouterTest(t, []string{"secure.example.com"})
manualResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/tls-certificates/", map[string]any{
"name": "secure-example",
"cert_pem": certPEM,
"key_pem": keyPEM,
"remark": "manual import",
})
var manualCertificate model.TLSCertificate
decodeResponseData(t, manualResp, &manualCertificate)
if manualCertificate.ID == 0 {
t.Fatal("expected manual certificate import to persist certificate")
}
detailResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/tls-certificates/"+toString(manualCertificate.ID), nil)
var certificateDetail map[string]any
decodeResponseData(t, detailResp, &certificateDetail)
if _, exists := certificateDetail["cert_pem"]; exists {
t.Fatal("expected certificate detail endpoint to omit cert_pem")
}
if _, exists := certificateDetail["key_pem"]; exists {
t.Fatal("expected certificate detail endpoint to omit key_pem")
}
contentResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/tls-certificates/"+toString(manualCertificate.ID)+"/content", nil)
var certificateContent map[string]any
decodeResponseData(t, contentResp, &certificateContent)
if certificateContent["cert_pem"] == "" || certificateContent["key_pem"] == "" {
t.Fatal("expected certificate content endpoint to return pem payloads")
}
updatedCertPEM, updatedKeyPEM := generateCertificatePairForRouterTest(t, []string{"secure.example.com", "www.secure.example.com"})
updateCertificateResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/tls-certificates/"+toString(manualCertificate.ID)+"/update", map[string]any{
"name": "secure-example-updated",
"cert_pem": updatedCertPEM,
"key_pem": updatedKeyPEM,
"remark": "updated manual import",
})
decodeResponseData(t, updateCertificateResp, &manualCertificate)
if manualCertificate.Name != "secure-example-updated" || manualCertificate.Remark != "updated manual import" {
t.Fatalf("expected certificate update to persist metadata, got %+v", manualCertificate)
}
fileCertPEM, fileKeyPEM := generateCertificatePairForRouterTest(t, []string{"upload.example.com"})
multipartResp := performMultipartRequest(t, engine, token, "/api/tls-certificates/import-file", map[string]string{
"name": "upload-example",
"remark": "upload import",
}, map[string]string{
"cert_file": fileCertPEM,
"key_file": fileKeyPEM,
})
var uploadedCertificate model.TLSCertificate
decodeResponseData(t, multipartResp, &uploadedCertificate)
if uploadedCertificate.ID == 0 {
t.Fatal("expected file certificate import to persist certificate")
}
resp := performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/", map[string]any{
"domain": "secure.example.com",
"origin_url": "https://origin-secure.internal",
"enabled": true,
"enable_https": true,
"cert_id": manualCertificate.ID,
"redirect_http": true,
"remark": "https route",
})
var route service.ProxyRouteView
decodeResponseData(t, resp, &route)
if !route.EnableHTTPS || route.CertID == nil || *route.CertID != manualCertificate.ID {
t.Fatal("expected route to persist https certificate binding")
}
updateResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/"+toString(route.ID)+"/update", map[string]any{
"domain": "secure.example.com",
"origin_url": "http://origin-secure.internal",
"enabled": true,
"enable_https": false,
"cert_id": nil,
"redirect_http": false,
"remark": "downgraded route",
})
decodeResponseData(t, updateResp, &route)
if route.EnableHTTPS || route.CertID != nil || route.RedirectHTTP {
t.Fatalf("expected route to disable https flags, got %+v", route)
}
updateResp = performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/"+toString(route.ID)+"/update", map[string]any{
"domain": "secure.example.com",
"origin_url": "https://origin-secure.internal",
"enabled": true,
"enable_https": true,
"cert_id": manualCertificate.ID,
"redirect_http": true,
"remark": "re-enabled https route",
})
decodeResponseData(t, updateResp, &route)
if !route.EnableHTTPS || route.CertID == nil || *route.CertID != manualCertificate.ID || !route.RedirectHTTP {
t.Fatalf("expected route update to persist https fields, got %+v", route)
}
listResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/proxy-routes/", nil)
var routes []service.ProxyRouteView
decodeResponseData(t, listResp, &routes)
if len(routes) != 1 || !routes[0].EnableHTTPS || routes[0].CertID == nil || *routes[0].CertID != manualCertificate.ID || !routes[0].RedirectHTTP {
t.Fatalf("expected route list to reflect https update, got %+v", routes)
}
certificateListResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/tls-certificates/", nil)
var certificateList []map[string]any
decodeResponseData(t, certificateListResp, &certificateList)
if len(certificateList) == 0 {
t.Fatal("expected certificate list to return records")
}
if _, exists := certificateList[0]["cert_pem"]; exists {
t.Fatal("expected certificate list to omit cert_pem")
}
if _, exists := certificateList[0]["key_pem"]; exists {
t.Fatal("expected certificate list to omit key_pem")
}
resp = performJSONRequest(t, engine, token, http.MethodPost, "/api/config-versions/publish", nil)
var version model.ConfigVersion
decodeResponseData(t, resp, &version)
if !strings.Contains(version.MainConfig, "include __OPENFLARE_ROUTE_CONFIG__;") {
t.Fatal("expected active config to render managed main config")
}
if !strings.Contains(version.RenderedConfig, "listen 443 ssl;") {
t.Fatal("expected active config to render https ssl listener")
}
if !strings.Contains(version.RenderedConfig, "http2 on;") {
t.Fatal("expected active config to render dedicated http2 directive")
}
if !strings.Contains(version.RenderedConfig, "return 301 https://$host$request_uri;") {
t.Fatal("expected active config to render redirect server")
}
if !strings.Contains(version.SupportFilesJSON, ".crt") || !strings.Contains(version.SupportFilesJSON, ".key") {
t.Fatal("expected support files json to contain certificate artifacts")
}
if err := (&model.Node{
NodeID: "phase1-node",
Name: "phase1-node",
IP: "10.0.0.8",
AccessToken: common.AccessToken,
Version: "0.1.0",
ExtVersion: "1.25.5",
Status: service.NodeStatusOnline,
LastSeenAt: time.Now(),
}).Insert(); err != nil {
t.Fatalf("failed to seed phase1 node: %v", err)
}
agentResp := performAgentJSONRequestWithToken(t, engine, common.AccessToken, http.MethodGet, "/api/agent/config-versions/active", nil)
var activeConfig map[string]any
decodeResponseData(t, agentResp, &activeConfig)
sourceConfigJSON, ok := activeConfig["source_config_json"].(string)
if !ok || !strings.Contains(sourceConfigJSON, "secure.example.com") {
t.Fatalf("expected active config to expose source_config_json, got %#v", activeConfig["source_config_json"])
}
supportFiles, ok := activeConfig["support_files"].([]any)
if !ok || len(supportFiles) != 2 {
t.Fatalf("expected active config to expose 2 certificate support files, got %#v", activeConfig["support_files"])
}
}
func TestTLSCertificateConvertAcmeAPI(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
token := prepareRootToken(t)
certPEM, keyPEM := generateCertificatePairForRouterTest(t, []string{"manual.example.com"})
createResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/tls-certificates/", map[string]any{
"name": "manual-example",
"cert_pem": certPEM,
"key_pem": keyPEM,
})
var certificate model.TLSCertificate
decodeResponseData(t, createResp, &certificate)
started := make(chan struct{}, 1)
release := make(chan struct{})
done := make(chan struct{})
restore := service.SetTLSCertificateObtainFuncForTest(func(c *model.TLSCertificate) error {
defer close(done)
started <- struct{}{}
<-release
return errors.New("stop test conversion before external ACME call")
})
t.Cleanup(func() {
close(release)
<-done
restore()
})
convertResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/tls-certificates/"+toString(certificate.ID)+"/convert-acme", map[string]any{
"name": "managed-example",
"remark": "convert via api",
"acme_account_id": 1,
"dns_account_id": 2,
"key_algorithm": "EC256",
"auto_renew": true,
"primary_domain": "manual.example.com",
})
var converted model.TLSCertificate
decodeResponseData(t, convertResp, &converted)
if converted.ID != certificate.ID || converted.Provider != "upload" || converted.ApplyStatus != "applying" {
t.Fatalf("expected conversion API to keep upload provider while applying, got %+v", converted)
}
select {
case <-started:
case <-time.After(time.Second):
t.Fatal("expected conversion task to start")
}
duplicateResp := performJSONRequestNoFatal(t, engine, token, http.MethodPost, "/api/tls-certificates/"+toString(certificate.ID)+"/convert-acme", map[string]any{
"name": "managed-example",
"primary_domain": "manual.example.com",
})
if duplicateResp.Success || !strings.Contains(duplicateResp.Message, "already applying") {
t.Fatalf("expected duplicate conversion to fail, got %+v", duplicateResp)
}
invalidResp := performJSONRequestNoFatal(t, engine, token, http.MethodPost, "/api/tls-certificates/not-a-number/convert-acme", map[string]any{})
if invalidResp.Success || !strings.Contains(invalidResp.Message, "参数错误") {
t.Fatalf("expected invalid id to fail, got %+v", invalidResp)
}
acmeCertPEM, acmeKeyPEM := generateCertificatePairForRouterTest(t, []string{"acme.example.com"})
acmeResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/tls-certificates/", map[string]any{
"name": "already-acme",
"cert_pem": acmeCertPEM,
"key_pem": acmeKeyPEM,
})
var acmeCertificate model.TLSCertificate
decodeResponseData(t, acmeResp, &acmeCertificate)
acmeCertificate.Provider = "acme"
if err := acmeCertificate.Update(); err != nil {
t.Fatalf("failed to mark certificate acme: %v", err)
}
nonUploadResp := performJSONRequestNoFatal(t, engine, token, http.MethodPost, "/api/tls-certificates/"+toString(acmeCertificate.ID)+"/convert-acme", map[string]any{
"name": "already-acme",
"primary_domain": "acme.example.com",
})
if nonUploadResp.Success || !strings.Contains(nonUploadResp.Message, "only uploaded") {
t.Fatalf("expected non-upload conversion to fail, got %+v", nonUploadResp)
}
}
func setupTestDB(t *testing.T) {
t.Helper()
dbPath := filepath.Join(t.TempDir(), "phase1.db")
common.SQLitePath = dbPath
common.AccessToken = "phase1-agent-token"
originalCapLoginEnabled := common.CapLoginEnabled
common.CapLoginEnabled = false
if err := model.InitDB(); err != nil {
t.Fatalf("failed to init db: %v", err)
}
middleware.InitJWTMiddleware()
t.Cleanup(func() {
common.CapLoginEnabled = originalCapLoginEnabled
if err := model.CloseDB(); err != nil {
t.Fatalf("failed to close db: %v", err)
}
})
}
func prepareRootToken(t *testing.T) string {
t.Helper()
user := &model.User{Username: "root"}
if err := user.FillUserByUsername(); err != nil {
t.Fatalf("failed to load root user: %v", err)
}
// Generate a proper JWT so auth middleware can validate it
tokenString, _, err := middleware.JWTMiddleware.TokenGenerator(user)
if err != nil {
t.Fatalf("failed to generate JWT for root user: %v", err)
}
if err := model.DB.Model(user).Update("token", tokenString).Error; err != nil {
t.Fatalf("failed to set root token: %v", err)
}
return tokenString
}
func performJSONRequest(t *testing.T, engine http.Handler, token string, method string, path string, body any) apiResponse {
t.Helper()
var payload []byte
var err error
if body != nil {
payload, err = json.Marshal(body)
if err != nil {
t.Fatalf("failed to marshal request body: %v", err)
}
}
req := httptest.NewRequest(method, path, bytes.NewReader(payload))
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
req.Header.Set("OpenFlare-Token", token)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected status %d for %s %s: %s", recorder.Code, method, path, recorder.Body.String())
}
var resp apiResponse
if err = json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if !resp.Success {
t.Fatalf("request %s %s failed: %s", method, path, resp.Message)
}
return resp
}
func performJSONRequestNoFatal(t *testing.T, engine http.Handler, token string, method string, path string, body any) apiResponse {
t.Helper()
var payload []byte
var err error
if body != nil {
payload, err = json.Marshal(body)
if err != nil {
t.Fatalf("failed to marshal request body: %v", err)
}
}
req := httptest.NewRequest(method, path, bytes.NewReader(payload))
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
req.Header.Set("OpenFlare-Token", token)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK && recorder.Code != http.StatusBadRequest {
t.Fatalf("unexpected status %d for %s %s: %s", recorder.Code, method, path, recorder.Body.String())
}
var resp apiResponse
if err = json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
return resp
}
func decodeResponseData(t *testing.T, resp apiResponse, target any) {
t.Helper()
if err := json.Unmarshal(resp.Data, target); err != nil {
t.Fatalf("failed to decode response data: %v", err)
}
}
func toString(id uint) string {
return strconv.FormatUint(uint64(id), 10)
}
func performMultipartRequest(t *testing.T, engine http.Handler, token string, path string, fields map[string]string, files map[string]string) apiResponse {
t.Helper()
var body bytes.Buffer
writer := multipart.NewWriter(&body)
for key, value := range fields {
if err := writer.WriteField(key, value); err != nil {
t.Fatalf("failed to write multipart field: %v", err)
}
}
for fieldName, content := range files {
part, err := writer.CreateFormFile(fieldName, fieldName+".pem")
if err != nil {
t.Fatalf("failed to create multipart file: %v", err)
}
if _, err = part.Write([]byte(content)); err != nil {
t.Fatalf("failed to write multipart file content: %v", err)
}
}
if err := writer.Close(); err != nil {
t.Fatalf("failed to close multipart writer: %v", err)
}
req := httptest.NewRequest(http.MethodPost, path, &body)
req.Header.Set("Content-Type", writer.FormDataContentType())
req.Header.Set("OpenFlare-Token", token)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected status %d for multipart %s: %s", recorder.Code, path, recorder.Body.String())
}
var resp apiResponse
if err := json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal multipart response: %v", err)
}
if !resp.Success {
t.Fatalf("multipart request %s failed: %s", path, resp.Message)
}
return resp
}
func generateCertificatePairForRouterTest(t *testing.T, dnsNames []string) (string, string) {
t.Helper()
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatalf("GenerateKey failed: %v", err)
}
template := &x509.Certificate{
SerialNumber: big.NewInt(time.Now().UnixNano()),
Subject: pkix.Name{
CommonName: dnsNames[0],
},
DNSNames: dnsNames,
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(24 * time.Hour),
KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
}
certDER, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey)
if err != nil {
t.Fatalf("CreateCertificate failed: %v", err)
}
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER})
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(privateKey)})
return string(certPEM), string(keyPEM)
}
@@ -0,0 +1,113 @@
package router_test
import (
"net/http"
"testing"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/router"
)
func TestPhase2ManagedDomainLifecycle(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
token := prepareRootToken(t)
wildcardCertPEM, wildcardKeyPEM := generateCertificatePairForRouterTest(t, []string{"*.example.com"})
exactCertPEM, exactKeyPEM := generateCertificatePairForRouterTest(t, []string{"api.example.com"})
wildcardResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/tls-certificates/", map[string]any{
"name": "wildcard-cert",
"cert_pem": wildcardCertPEM,
"key_pem": wildcardKeyPEM,
})
var wildcardCertificate map[string]any
decodeResponseData(t, wildcardResp, &wildcardCertificate)
exactResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/tls-certificates/", map[string]any{
"name": "exact-cert",
"cert_pem": exactCertPEM,
"key_pem": exactKeyPEM,
})
var exactCertificate map[string]any
decodeResponseData(t, exactResp, &exactCertificate)
wildcardID := uint(wildcardCertificate["id"].(float64))
exactID := uint(exactCertificate["id"].(float64))
createWildcard := performJSONRequest(t, engine, token, http.MethodPost, "/api/managed-domains/", map[string]any{
"domain": "*.example.com",
"cert_id": wildcardID,
"enabled": true,
"remark": "wildcard binding",
})
var wildcardDomain map[string]any
decodeResponseData(t, createWildcard, &wildcardDomain)
createExact := performJSONRequest(t, engine, token, http.MethodPost, "/api/managed-domains/", map[string]any{
"domain": "api.example.com",
"cert_id": exactID,
"enabled": true,
"remark": "exact binding",
})
var exactDomain map[string]any
decodeResponseData(t, createExact, &exactDomain)
listResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/managed-domains/", nil)
var domains []map[string]any
decodeResponseData(t, listResp, &domains)
if len(domains) != 2 {
t.Fatalf("expected 2 managed domains, got %d", len(domains))
}
matchResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/managed-domains/match?domain=api.example.com", nil)
var matchResult map[string]any
decodeResponseData(t, matchResp, &matchResult)
if matched, ok := matchResult["matched"].(bool); !ok || !matched {
t.Fatalf("expected exact domain to be matched, got %#v", matchResult)
}
candidate, ok := matchResult["candidate"].(map[string]any)
if !ok {
t.Fatalf("expected candidate payload, got %#v", matchResult["candidate"])
}
if candidate["match_type"] != "exact" {
t.Fatalf("expected exact match type, got %#v", candidate["match_type"])
}
if uint(candidate["certificate_id"].(float64)) != exactID {
t.Fatalf("expected exact certificate id %d, got %#v", exactID, candidate["certificate_id"])
}
updateResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/managed-domains/"+toString(uint(exactDomain["id"].(float64)))+"/update", map[string]any{
"domain": "api.example.com",
"cert_id": exactID,
"enabled": false,
"remark": "disabled exact binding",
})
decodeResponseData(t, updateResp, &exactDomain)
matchResp = performJSONRequest(t, engine, token, http.MethodGet, "/api/managed-domains/match?domain=api.example.com", nil)
decodeResponseData(t, matchResp, &matchResult)
candidate, ok = matchResult["candidate"].(map[string]any)
if !ok {
t.Fatalf("expected wildcard fallback candidate, got %#v", matchResult["candidate"])
}
if candidate["match_type"] != "wildcard" {
t.Fatalf("expected wildcard fallback, got %#v", candidate["match_type"])
}
if uint(candidate["certificate_id"].(float64)) != wildcardID {
t.Fatalf("expected wildcard certificate id %d, got %#v", wildcardID, candidate["certificate_id"])
}
deleteResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/managed-domains/"+toString(uint(wildcardDomain["id"].(float64)))+"/delete", nil)
if !deleteResp.Success {
t.Fatalf("expected delete success, got %s", deleteResp.Message)
}
}
@@ -0,0 +1,930 @@
package router_test
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/rain-kl/openflare/openflare-server/internal/router"
"github.com/rain-kl/openflare/openflare-server/internal/service"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
)
func TestPhase2RateLimitOptionsHotReload(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
model.InitOptionMap()
oldGlobalApiRateLimitNum := common.GlobalApiRateLimitNum
oldGlobalApiRateLimitDuration := common.GlobalApiRateLimitDuration
oldCriticalRateLimitNum := common.CriticalRateLimitNum
oldCriticalRateLimitDuration := common.CriticalRateLimitDuration
t.Cleanup(func() {
common.GlobalApiRateLimitNum = oldGlobalApiRateLimitNum
common.GlobalApiRateLimitDuration = oldGlobalApiRateLimitDuration
common.CriticalRateLimitNum = oldCriticalRateLimitNum
common.CriticalRateLimitDuration = oldCriticalRateLimitDuration
})
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginCookie := loginAsRoot(t, engine)
performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/option/update-batch", map[string]any{
"options": []map[string]any{
{
"key": "GlobalApiRateLimitNum",
"value": "450",
},
{
"key": "GlobalApiRateLimitDuration",
"value": "240",
},
{
"key": "CriticalRateLimitNum",
"value": "150",
},
{
"key": "CriticalRateLimitDuration",
"value": "900",
},
},
})
if common.GlobalApiRateLimitNum != 450 {
t.Fatalf("expected GlobalApiRateLimitNum to be hot reloaded, got %d", common.GlobalApiRateLimitNum)
}
if common.GlobalApiRateLimitDuration != 240 {
t.Fatalf("expected GlobalApiRateLimitDuration to be hot reloaded, got %d", common.GlobalApiRateLimitDuration)
}
if common.CriticalRateLimitNum != 150 {
t.Fatalf("expected CriticalRateLimitNum to be hot reloaded, got %d", common.CriticalRateLimitNum)
}
if common.CriticalRateLimitDuration != 900 {
t.Fatalf("expected CriticalRateLimitDuration to be hot reloaded, got %d", common.CriticalRateLimitDuration)
}
resp := performSessionJSONRequest(t, engine, loginCookie, http.MethodGet, "/api/option/", nil)
var options []model.Option
decodeResponseData(t, resp, &options)
optionMap := make(map[string]string, len(options))
for _, option := range options {
optionMap[option.Key] = option.Value
}
if optionMap["GlobalApiRateLimitNum"] != "450" {
t.Fatalf("expected option payload to include GlobalApiRateLimitNum=450, got %q", optionMap["GlobalApiRateLimitNum"])
}
if optionMap["CriticalRateLimitDuration"] != "900" {
t.Fatalf("expected option payload to include CriticalRateLimitDuration=900, got %q", optionMap["CriticalRateLimitDuration"])
}
}
func TestPhase2BatchOptionUpdateIsAtomic(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
model.InitOptionMap()
oldGlobalAPI := common.GlobalApiRateLimitNum
t.Cleanup(func() {
common.GlobalApiRateLimitNum = oldGlobalAPI
})
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginCookie := loginAsRoot(t, engine)
payload, err := json.Marshal(map[string]any{
"options": []map[string]any{
{
"key": "GlobalApiRateLimitNum",
"value": "451",
},
{
"key": "CriticalRateLimitDuration",
"value": "1800",
},
},
})
if err != nil {
t.Fatalf("failed to marshal batch payload: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/api/option/update-batch", bytes.NewReader(payload))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("OpenFlare-Token", loginCookie)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected status %d: %s", recorder.Code, recorder.Body.String())
}
var resp apiResponse
if err = json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if resp.Success {
t.Fatal("expected invalid batch update to fail")
}
if common.GlobalApiRateLimitNum != oldGlobalAPI {
t.Fatalf("expected GlobalApiRateLimitNum to remain %d after failed batch, got %d", oldGlobalAPI, common.GlobalApiRateLimitNum)
}
resp = performSessionJSONRequest(t, engine, loginCookie, http.MethodGet, "/api/option/", nil)
var options []model.Option
decodeResponseData(t, resp, &options)
optionMap := make(map[string]string, len(options))
for _, option := range options {
optionMap[option.Key] = option.Value
}
if optionMap["GlobalApiRateLimitNum"] == "451" {
t.Fatal("expected failed batch update to avoid persisting partial values")
}
}
func TestPhase2BatchOptionUpdateValidatesMergedState(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
model.InitOptionMap()
oldGitHubClientID := common.GitHubClientId
oldGitHubOAuthEnabled := common.GitHubOAuthEnabled
t.Cleanup(func() {
common.GitHubClientId = oldGitHubClientID
common.GitHubOAuthEnabled = oldGitHubOAuthEnabled
})
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginCookie := loginAsRoot(t, engine)
performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/option/update-batch", map[string]any{
"options": []map[string]any{
{
"key": "GitHubClientId",
"value": "client-id-from-batch",
},
{
"key": "GitHubOAuthEnabled",
"value": "true",
},
},
})
if common.GitHubClientId != "client-id-from-batch" {
t.Fatalf("expected GitHubClientId to be updated from batch, got %q", common.GitHubClientId)
}
if !common.GitHubOAuthEnabled {
t.Fatal("expected GitHubOAuthEnabled to be enabled by merged batch state")
}
}
func TestAuthSourceUpdateAcceptsClientSecret(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginCookie := loginAsRoot(t, engine)
createResp := performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/auth-sources/", map[string]any{
"name": "GitHub",
"type": "github",
"display_name": "GitHub",
"is_active": false,
"client_id": "github-client-id",
"client_secret": "initial-secret",
"scopes": "user:email",
})
var created model.AuthSource
decodeResponseData(t, createResp, &created)
if created.ClientSecret != "" {
t.Fatal("expected create response to avoid exposing client_secret")
}
if !created.ClientSecretConfigured {
t.Fatal("expected create response to mark client_secret as configured")
}
updateResp := performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/auth-sources/1/update", map[string]any{
"name": "GitHub",
"type": "github",
"display_name": "GitHub",
"is_active": true,
"client_id": "github-client-id",
"client_secret": "updated-secret",
"scopes": "user:email",
})
var updated model.AuthSource
decodeResponseData(t, updateResp, &updated)
if updated.ClientSecret != "" {
t.Fatal("expected update response to avoid exposing client_secret")
}
if !updated.ClientSecretConfigured {
t.Fatal("expected update response to mark client_secret as configured")
}
if !updated.IsActive {
t.Fatal("expected auth source to be active after update")
}
stored, err := model.GetAuthSourceByID(1)
if err != nil {
t.Fatalf("expected auth source to exist: %v", err)
}
if stored.ClientSecret != "updated-secret" {
t.Fatalf("expected stored client secret to be updated, got %q", stored.ClientSecret)
}
performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/auth-sources/1/toggle", map[string]any{
"is_active": false,
})
performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/auth-sources/1/toggle", map[string]any{
"is_active": true,
})
}
func TestExternalAccountBindingsCanBeListedAndDeleted(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginCookie := loginAsRoot(t, engine)
source := &model.AuthSource{
Name: "logto",
Type: model.AuthSourceTypeOIDC,
DisplayName: "Logto",
ClientID: "logto-client-id",
ClientSecret: "logto-client-secret",
OpenIDDiscoveryURL: "https://auth.example.com/.well-known/openid-configuration",
}
if err := model.CreateAuthSource(source); err != nil {
t.Fatalf("create auth source: %v", err)
}
if err := model.LinkExternalAccount(&model.ExternalAccount{
AuthSourceID: source.ID,
UserID: 1,
ExternalID: "logto-user-1",
ExternalUsername: "ryan",
Email: "ryan@example.com",
}); err != nil {
t.Fatalf("link external account: %v", err)
}
listResp := performSessionJSONRequest(t, engine, loginCookie, http.MethodGet, "/api/oauth/external-accounts/", nil)
var bindings []model.ExternalAccountView
decodeResponseData(t, listResp, &bindings)
if len(bindings) != 1 {
t.Fatalf("expected 1 binding, got %d", len(bindings))
}
if bindings[0].AuthSourceName != "logto" || bindings[0].ExternalUsername != "ryan" {
t.Fatalf("unexpected binding view: %+v", bindings[0])
}
performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/oauth/external-accounts/1/delete", nil)
listResp = performSessionJSONRequest(t, engine, loginCookie, http.MethodGet, "/api/oauth/external-accounts/", nil)
decodeResponseData(t, listResp, &bindings)
if len(bindings) != 0 {
t.Fatalf("expected binding to be deleted, got %+v", bindings)
}
}
func loginAsRoot(t *testing.T, engine http.Handler) string {
t.Helper()
payload, err := json.Marshal(map[string]any{
"username": "root",
"password": "123456",
})
if err != nil {
t.Fatalf("failed to marshal login payload: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/api/user/login", bytes.NewReader(payload))
req.Header.Set("Content-Type", "application/json")
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected login status %d: %s", recorder.Code, recorder.Body.String())
}
var resp apiResponse
if err = json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to decode login response: %v", err)
}
if !resp.Success {
t.Fatalf("root login failed: %s", resp.Message)
}
var user model.User
if err = json.Unmarshal(resp.Data, &user); err != nil {
t.Fatalf("failed to decode login user: %v", err)
}
if user.Token == "" {
t.Fatal("expected OpenFlare-Token after root login")
}
return user.Token
}
func performSessionJSONRequest(t *testing.T, engine http.Handler, token string, method string, path string, body any) apiResponse {
t.Helper()
var payload []byte
var err error
if body != nil {
payload, err = json.Marshal(body)
if err != nil {
t.Fatalf("failed to marshal request body: %v", err)
}
}
req := httptest.NewRequest(method, path, bytes.NewReader(payload))
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
req.Header.Set("OpenFlare-Token", token)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected status %d for %s %s: %s", recorder.Code, method, path, recorder.Body.String())
}
var resp apiResponse
if err = json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if !resp.Success {
t.Fatalf("request %s %s failed: %s", method, path, resp.Message)
}
return resp
}
func TestPhase2AgentLifecycle(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
adminToken := prepareRootToken(t)
createRouteAndPublishVersion(t, engine, adminToken)
dashboardResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/dashboard/overview", nil)
var dashboard struct {
Summary service.DashboardSummary `json:"summary"`
}
decodeResponseData(t, dashboardResp, &dashboard)
if dashboard.Summary.TotalNodes != 0 {
t.Fatalf("expected empty dashboard node summary before node registration, got %+v", dashboard.Summary)
}
unauthorizedRequest := httptest.NewRequest(http.MethodPost, "/api/agent/nodes/register", bytes.NewReader([]byte(`{}`)))
unauthorizedRecorder := httptest.NewRecorder()
engine.ServeHTTP(unauthorizedRecorder, unauthorizedRequest)
if unauthorizedRecorder.Code != http.StatusUnauthorized {
t.Fatalf("expected unauthorized status for missing discovery token, got %d", unauthorizedRecorder.Code)
}
createdNodeResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/nodes/", map[string]any{
"name": "shanghai-edge-1",
"geo_manual_override": true,
"geo_name": "Shanghai",
"geo_latitude": 31.2304,
"geo_longitude": 121.4737,
})
var createdNode service.NodeView
decodeResponseData(t, createdNodeResp, &createdNode)
if createdNode.AccessToken == "" || createdNode.Status != service.NodeStatusPending {
t.Fatal("expected created node to expose agent token with pending status")
}
if createdNode.GeoName != "Shanghai" || createdNode.GeoLatitude == nil || createdNode.GeoLongitude == nil {
t.Fatalf("expected created node to expose geo metadata, got %+v", createdNode)
}
heartbeatPayload := map[string]any{
"node_id": "spoofed-node-id",
"name": "shanghai-edge-1",
"ip": "10.0.0.9",
"version": "0.1.1",
"ext_version": "1.27.1.2",
"openresty_status": service.OpenrestyStatusUnhealthy,
"openresty_message": "docker run openresty failed: bind 80 already allocated",
"current_version": "",
"last_error": "",
}
resp := performAgentJSONRequestWithTokenAndRemote(t, engine, createdNode.AccessToken, http.MethodPost, "/api/agent/nodes/heartbeat", heartbeatPayload, "198.51.100.10:1234")
var registeredNode model.Node
decodeResponseData(t, resp, &registeredNode)
if registeredNode.IP != "198.51.100.10" || registeredNode.Version != "0.1.1" || registeredNode.NodeID != createdNode.NodeID {
t.Fatal("expected heartbeat to update node metadata")
}
if registeredNode.OpenrestyStatus != service.OpenrestyStatusUnhealthy {
t.Fatal("expected heartbeat to update openresty status")
}
activeConfigResp := performAgentJSONRequestWithToken(t, engine, createdNode.AccessToken, http.MethodGet, "/api/agent/config-versions/active", nil)
var activeConfig service.AgentConfigResponse
decodeResponseData(t, activeConfigResp, &activeConfig)
if activeConfig.Version == "" || activeConfig.SourceConfigJSON == "" || activeConfig.Checksum == "" {
t.Fatal("expected active config response to contain version payload")
}
successApplyResp := performAgentJSONRequestWithToken(t, engine, createdNode.AccessToken, http.MethodPost, "/api/agent/apply-logs", map[string]any{
"node_id": "spoofed-node-id",
"version": activeConfig.Version,
"result": service.ApplyResultOK,
"message": "apply ok",
})
var successApplyLog model.ApplyLog
decodeResponseData(t, successApplyResp, &successApplyLog)
if successApplyLog.Result != service.ApplyResultOK {
t.Fatal("expected apply log success to be recorded")
}
failedApplyResp := performAgentJSONRequestWithToken(t, engine, createdNode.AccessToken, http.MethodPost, "/api/agent/apply-logs", map[string]any{
"node_id": "spoofed-node-id",
"version": activeConfig.Version,
"result": service.ApplyResultFailed,
"message": "openresty reload failed",
})
var failedApplyLog model.ApplyLog
decodeResponseData(t, failedApplyResp, &failedApplyLog)
if failedApplyLog.Result != service.ApplyResultFailed {
t.Fatal("expected failed apply log to be recorded")
}
nodesResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/nodes/", nil)
var nodes []service.NodeView
decodeResponseData(t, nodesResp, &nodes)
if len(nodes) != 1 {
t.Fatalf("expected 1 node, got %d", len(nodes))
}
if nodes[0].Status != service.NodeStatusOnline {
t.Fatal("expected registered node to become online")
}
if nodes[0].AccessToken != createdNode.AccessToken {
t.Fatal("expected node auth token to remain stable after occupancy")
}
if nodes[0].LatestApplyResult != service.ApplyResultFailed || nodes[0].LatestApplyMessage != "openresty reload failed" {
t.Fatal("expected node list to expose latest apply status")
}
if nodes[0].CurrentVersion != activeConfig.Version {
t.Fatal("expected node current_version to remain at last successful version")
}
if nodes[0].LastError != "openresty reload failed" {
t.Fatal("expected node last_error to reflect failed apply")
}
if nodes[0].OpenrestyStatus != service.OpenrestyStatusUnhealthy {
t.Fatal("expected node list to expose openresty status")
}
if nodes[0].OpenrestyMessage != "docker run openresty failed: bind 80 already allocated" {
t.Fatal("expected node list to expose openresty message")
}
if err := model.DB.Create(&model.NodeHealthEvent{
NodeID: createdNode.NodeID,
EventType: "openresty_down",
Severity: service.NodeHealthSeverityCritical,
Status: service.NodeHealthEventStatusActive,
Message: "docker run openresty failed: bind 80 already allocated",
FirstTriggeredAt: time.Now().Add(-2 * time.Minute),
LastTriggeredAt: time.Now().Add(-time.Minute),
ReportedAt: time.Now().Add(-time.Minute),
}).Error; err != nil {
t.Fatalf("failed to insert node health event: %v", err)
}
observabilityResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/nodes/"+toString(createdNode.ID)+"/observability?hours=24&limit=20", nil)
var observability service.NodeObservabilityView
decodeResponseData(t, observabilityResp, &observability)
if observability.NodeID != createdNode.NodeID {
t.Fatalf("expected observability response for node %s, got %s", createdNode.NodeID, observability.NodeID)
}
if len(observability.HealthEvents) != 1 {
t.Fatalf("expected observability response to include health events, got %+v", observability.HealthEvents)
}
cleanupHealthResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/nodes/"+toString(createdNode.ID)+"/observability/cleanup", nil)
var cleanupHealthResult service.NodeHealthEventCleanupResult
decodeResponseData(t, cleanupHealthResp, &cleanupHealthResult)
if cleanupHealthResult.NodeID != createdNode.NodeID || cleanupHealthResult.DeletedCount != 1 {
t.Fatalf("unexpected node health cleanup result: %+v", cleanupHealthResult)
}
observabilityAfterCleanupResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/nodes/"+toString(createdNode.ID)+"/observability?hours=24&limit=20", nil)
decodeResponseData(t, observabilityAfterCleanupResp, &observability)
if len(observability.HealthEvents) != 0 {
t.Fatalf("expected health events to be cleaned up, got %+v", observability.HealthEvents)
}
restartResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/nodes/"+toString(createdNode.ID)+"/openresty-restart", nil)
decodeResponseData(t, restartResp, &createdNode)
if !createdNode.RestartOpenrestyRequested {
t.Fatal("expected openresty restart request flag to be set")
}
rawHeartbeatPayload, err := json.Marshal(heartbeatPayload)
if err != nil {
t.Fatalf("failed to marshal heartbeat payload: %v", err)
}
restartHeartbeatReq := httptest.NewRequest(http.MethodPost, "/api/agent/nodes/heartbeat", bytes.NewReader(rawHeartbeatPayload))
restartHeartbeatReq.Header.Set("Content-Type", "application/json")
restartHeartbeatReq.Header.Set("X-Agent-Token", createdNode.AccessToken)
restartHeartbeatReq.RemoteAddr = "198.51.100.10:1234"
restartHeartbeatRecorder := httptest.NewRecorder()
engine.ServeHTTP(restartHeartbeatRecorder, restartHeartbeatReq)
if restartHeartbeatRecorder.Code != http.StatusOK {
t.Fatalf("unexpected heartbeat status %d: %s", restartHeartbeatRecorder.Code, restartHeartbeatRecorder.Body.String())
}
var restartHeartbeatBody struct {
Success bool `json:"success"`
Message string `json:"message"`
AgentSettings service.AgentSettings `json:"agent_settings"`
ActiveConfig *service.ActiveConfigMeta `json:"active_config"`
}
if err = json.Unmarshal(restartHeartbeatRecorder.Body.Bytes(), &restartHeartbeatBody); err != nil {
t.Fatalf("failed to decode heartbeat response: %v", err)
}
if !restartHeartbeatBody.Success {
t.Fatalf("expected heartbeat request success, got %s", restartHeartbeatBody.Message)
}
if !restartHeartbeatBody.AgentSettings.RestartOpenrestyNow {
t.Fatal("expected heartbeat response to instruct openresty restart")
}
if restartHeartbeatBody.ActiveConfig == nil || restartHeartbeatBody.ActiveConfig.Version == "" || restartHeartbeatBody.ActiveConfig.Checksum == "" {
t.Fatal("expected heartbeat response to include active config summary")
}
logsResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/apply-logs/?node_id="+createdNode.NodeID+"&pageNo=1&pageSize=1", nil)
var logs service.ApplyLogListResult
decodeResponseData(t, logsResp, &logs)
if logs.Current != 1 || logs.Total != 2 || logs.TotalPage != 2 {
t.Fatalf("unexpected paged apply logs result: %+v", logs)
}
if len(logs.Rows) != 1 {
t.Fatalf("expected 1 apply log row on page 1, got %d", len(logs.Rows))
}
if logs.Rows[0].Result != service.ApplyResultFailed {
t.Fatalf("expected newest apply log first, got %s", logs.Rows[0].Result)
}
oldApplyLogTime := time.Now().Add(-48 * time.Hour)
if err := model.DB.Model(&model.ApplyLog{}).Where("id = ?", successApplyLog.ID).Update("created_at", oldApplyLogTime).Error; err != nil {
t.Fatalf("failed to backdate apply log: %v", err)
}
cleanupResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/apply-logs/cleanup", map[string]any{
"retention_days": 1,
})
var cleanupResult service.ApplyLogCleanupResult
decodeResponseData(t, cleanupResp, &cleanupResult)
if cleanupResult.DeleteAll {
t.Fatal("expected retention cleanup instead of delete-all cleanup")
}
if cleanupResult.RetentionDays != 1 || cleanupResult.DeletedCount != 1 {
t.Fatalf("unexpected cleanup result: %+v", cleanupResult)
}
postCleanupResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/apply-logs/?node_id="+createdNode.NodeID, nil)
decodeResponseData(t, postCleanupResp, &logs)
if logs.Total != 1 || len(logs.Rows) != 1 {
t.Fatalf("expected one apply log after retention cleanup, got %+v", logs)
}
deleteAllResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/apply-logs/cleanup", map[string]any{
"delete_all": true,
})
decodeResponseData(t, deleteAllResp, &cleanupResult)
if !cleanupResult.DeleteAll || cleanupResult.DeletedCount != 1 {
t.Fatalf("unexpected delete-all cleanup result: %+v", cleanupResult)
}
emptyLogsResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/apply-logs/?node_id="+createdNode.NodeID, nil)
decodeResponseData(t, emptyLogsResp, &logs)
if logs.Total != 0 || len(logs.Rows) != 0 || logs.Current != 1 || logs.TotalPage != 0 {
t.Fatalf("expected empty apply log page after delete-all cleanup, got %+v", logs)
}
postDeleteApplyResp := performAgentJSONRequestWithToken(t, engine, createdNode.AccessToken, http.MethodPost, "/api/agent/apply-logs", map[string]any{
"version": activeConfig.Version,
"result": service.ApplyResultOK,
"message": "local config already matches active version; apply skipped",
"checksum": activeConfig.Checksum,
})
var postDeleteApplyLog model.ApplyLog
decodeResponseData(t, postDeleteApplyResp, &postDeleteApplyLog)
if postDeleteApplyLog.ID == 0 || postDeleteApplyLog.NodeID != createdNode.NodeID {
t.Fatalf("expected apply log to be recreated after delete-all cleanup, got %+v", postDeleteApplyLog)
}
postDeleteLogsResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/apply-logs/?node_id="+createdNode.NodeID, nil)
decodeResponseData(t, postDeleteLogsResp, &logs)
if logs.Total != 1 || len(logs.Rows) != 1 || logs.Rows[0].ID != postDeleteApplyLog.ID {
t.Fatalf("expected new apply log after delete-all cleanup, got %+v", logs)
}
updatedNodeResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/nodes/"+toString(createdNode.ID)+"/update", map[string]any{
"name": "shanghai-edge-1-renamed",
"geo_manual_override": true,
"geo_name": "Tokyo",
"geo_latitude": 35.6762,
"geo_longitude": 139.6503,
})
decodeResponseData(t, updatedNodeResp, &createdNode)
if createdNode.Name != "shanghai-edge-1-renamed" {
t.Fatal("expected node name to be editable")
}
if createdNode.GeoName != "Tokyo" || createdNode.GeoLatitude == nil || createdNode.GeoLongitude == nil {
t.Fatalf("expected node geo metadata to be editable, got %+v", createdNode)
}
oldTime := time.Now().Add(-common.NodeOfflineThreshold - time.Minute)
if err := model.DB.Model(&model.Node{}).Where("node_id = ?", createdNode.NodeID).Update("last_seen_at", oldTime).Error; err != nil {
t.Fatalf("failed to update node last_seen_at: %v", err)
}
nodesResp = performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/nodes/", nil)
decodeResponseData(t, nodesResp, &nodes)
if nodes[0].Status != service.NodeStatusOffline {
t.Fatal("expected node to be shown as offline after timeout")
}
deleteResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/nodes/"+toString(createdNode.ID)+"/delete", nil)
if !deleteResp.Success {
t.Fatalf("expected delete node success, got %s", deleteResp.Message)
}
deniedReq := httptest.NewRequest(http.MethodPost, "/api/agent/nodes/heartbeat", bytes.NewReader([]byte(`{"ip":"10.0.0.9","version":"0.1.1"}`)))
deniedReq.Header.Set("Content-Type", "application/json")
deniedReq.Header.Set("X-Agent-Token", createdNode.AccessToken)
deniedRecorder := httptest.NewRecorder()
engine.ServeHTTP(deniedRecorder, deniedReq)
if deniedRecorder.Code != http.StatusUnauthorized {
t.Fatalf("expected deleted node token to be rejected, got %d", deniedRecorder.Code)
}
}
func TestPhase2CustomHeadersPreviewAndDiffLifecycle(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
token := prepareRootToken(t)
createResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/", map[string]any{
"domain": "preview.example.com",
"origin_url": "https://origin-a.internal",
"origin_host": "preview-origin.internal",
"enabled": true,
"custom_headers": []map[string]any{
{"key": "X-Trace-Id", "value": "$request_id"},
},
})
var createdRoute service.ProxyRouteView
decodeResponseData(t, createResp, &createdRoute)
if !strings.Contains(createdRoute.CustomHeaders, "X-Trace-Id") {
t.Fatalf("expected custom headers to be stored as json, got %s", createdRoute.CustomHeaders)
}
if createdRoute.OriginHost != "preview-origin.internal" {
t.Fatalf("expected origin_host to be stored, got %s", createdRoute.OriginHost)
}
if createdRoute.SiteName != "preview.example.com" || createdRoute.PrimaryDomain != "preview.example.com" || createdRoute.DomainCount != 1 {
t.Fatalf("expected website identity fields in create response, got %+v", createdRoute)
}
performJSONRequest(t, engine, token, http.MethodPost, "/api/config-versions/publish", nil)
performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/"+toString(createdRoute.ID)+"/update", map[string]any{
"domain": "preview.example.com",
"origin_url": "https://origin-b.internal",
"origin_host": "preview-upstream.internal",
"enabled": true,
"custom_headers": []map[string]any{
{"key": "X-Trace-Id", "value": "$request_id"},
{"key": "X-Release", "value": "candidate"},
},
})
performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/", map[string]any{
"domain": "new-preview.example.com",
"origin_url": "https://origin-new.internal",
"enabled": true,
})
previewResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/config-versions/preview", nil)
var preview map[string]any
decodeResponseData(t, previewResp, &preview)
renderedConfig, _ := preview["rendered_config"].(string)
if websiteCount, ok := preview["website_count"].(float64); !ok || int(websiteCount) != 2 {
t.Fatalf("expected preview website_count=2, got %#v", preview["website_count"])
}
if !strings.Contains(renderedConfig, `proxy_set_header X-Release "candidate";`) {
t.Fatalf("expected preview endpoint to return custom header, got %s", renderedConfig)
}
if !strings.Contains(renderedConfig, `proxy_set_header Host "preview-upstream.internal";`) {
t.Fatalf("expected preview endpoint to return overridden host header, got %s", renderedConfig)
}
if !strings.Contains(renderedConfig, "proxy_ssl_server_name on;") {
t.Fatalf("expected preview endpoint to enable proxy ssl server name, got %s", renderedConfig)
}
if !strings.Contains(renderedConfig, `proxy_ssl_name "preview-upstream.internal";`) {
t.Fatalf("expected preview endpoint to return proxy ssl name, got %s", renderedConfig)
}
diffResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/config-versions/diff", nil)
var diff map[string]any
decodeResponseData(t, diffResp, &diff)
modifiedDomains, ok := diff["modified_domains"].([]any)
if !ok || len(modifiedDomains) != 1 || modifiedDomains[0].(string) != "preview.example.com" {
t.Fatalf("unexpected modified domains: %#v", diff["modified_domains"])
}
addedDomains, ok := diff["added_domains"].([]any)
if !ok || len(addedDomains) != 1 || addedDomains[0].(string) != "new-preview.example.com" {
t.Fatalf("unexpected added domains: %#v", diff["added_domains"])
}
modifiedSites, ok := diff["modified_sites"].([]any)
if !ok || len(modifiedSites) != 1 || modifiedSites[0].(string) != "preview.example.com" {
t.Fatalf("unexpected modified sites: %#v", diff["modified_sites"])
}
addedSites, ok := diff["added_sites"].([]any)
if !ok || len(addedSites) != 1 || addedSites[0].(string) != "new-preview.example.com" {
t.Fatalf("unexpected added sites: %#v", diff["added_sites"])
}
}
func TestPhase2ProxyRouteWebsiteDetailAndLimits(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
token := prepareRootToken(t)
createResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/", map[string]any{
"site_name": "marketing-site",
"domains": []string{"app.example.com", "www.example.com"},
"origin_url": "https://origin.internal",
"enabled": true,
"limit_conn_per_server": 120,
"limit_conn_per_ip": 12,
"limit_rate": "512K",
})
var createdRoute service.ProxyRouteView
decodeResponseData(t, createResp, &createdRoute)
if createdRoute.SiteName != "marketing-site" || createdRoute.PrimaryDomain != "app.example.com" {
t.Fatalf("unexpected create payload: %+v", createdRoute)
}
if createdRoute.DomainCount != 2 || len(createdRoute.Domains) != 2 || createdRoute.Domains[1] != "www.example.com" {
t.Fatalf("expected multi-domain website view, got %+v", createdRoute)
}
if createdRoute.LimitConnPerServer != 120 || createdRoute.LimitConnPerIP != 12 || createdRoute.LimitRate != "512k" {
t.Fatalf("expected normalized rate limit fields, got %+v", createdRoute)
}
if len(createdRoute.UpstreamList) != 1 || createdRoute.UpstreamList[0] != "https://origin.internal" {
t.Fatalf("expected structured upstream list, got %+v", createdRoute.UpstreamList)
}
detailResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/proxy-routes/"+toString(createdRoute.ID), nil)
var detail service.ProxyRouteView
decodeResponseData(t, detailResp, &detail)
if detail.ID != createdRoute.ID || detail.SiteName != "marketing-site" || detail.LimitRate != "512k" {
t.Fatalf("unexpected detail response: %+v", detail)
}
if len(detail.Domains) != 2 || detail.Domains[0] != "app.example.com" || detail.Domains[1] != "www.example.com" {
t.Fatalf("expected detail response to expose full domain list, got %+v", detail.Domains)
}
listResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/proxy-routes/", nil)
var routes []service.ProxyRouteView
decodeResponseData(t, listResp, &routes)
if len(routes) != 1 || routes[0].SiteName != "marketing-site" || routes[0].LimitConnPerServer != 120 {
t.Fatalf("unexpected proxy route list response: %+v", routes)
}
}
func TestPhase2GlobalDiscoveryRegistration(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
adminToken := prepareRootToken(t)
bootstrapResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/nodes/bootstrap-token", nil)
var bootstrap service.NodeBootstrapView
decodeResponseData(t, bootstrapResp, &bootstrap)
if bootstrap.DiscoveryToken == "" {
t.Fatal("expected global discovery token to be available")
}
resp := performAgentJSONRequestWithTokenAndRemote(t, engine, bootstrap.DiscoveryToken, http.MethodPost, "/api/agent/nodes/register", map[string]any{
"node_id": "local-node-id",
"name": "bulk-edge-1",
"ip": "10.0.0.18",
"version": "0.2.0",
"ext_version": "1.25.5",
"current_version": "",
"last_error": "",
}, "203.0.113.18:4321")
var registration service.AgentRegistrationResponse
decodeResponseData(t, resp, &registration)
if registration.AccessToken == "" || registration.NodeID == "" {
t.Fatal("expected discovery registration to issue node-specific agent token")
}
nodesResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/nodes/", nil)
var nodes []service.NodeView
decodeResponseData(t, nodesResp, &nodes)
if len(nodes) != 1 {
t.Fatalf("expected 1 discovered node, got %d", len(nodes))
}
if nodes[0].Name != "bulk-edge-1" || nodes[0].AccessToken != registration.AccessToken || nodes[0].Status != service.NodeStatusOnline {
t.Fatal("expected discovered node to be created online with issued agent token")
}
if nodes[0].IP != "203.0.113.18" {
t.Fatalf("expected discovered node to keep public source ip, got %s", nodes[0].IP)
}
}
func performAgentJSONRequestWithToken(t *testing.T, engine http.Handler, token string, method string, path string, body any) apiResponse {
return performAgentJSONRequestWithTokenAndRemote(t, engine, token, method, path, body, "")
}
func performAgentJSONRequestWithTokenAndRemote(t *testing.T, engine http.Handler, token string, method string, path string, body any, remoteAddr string) apiResponse {
t.Helper()
var payload []byte
var err error
if body != nil {
payload, err = json.Marshal(body)
if err != nil {
t.Fatalf("failed to marshal request body: %v", err)
}
}
req := httptest.NewRequest(method, path, bytes.NewReader(payload))
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
if remoteAddr != "" {
req.RemoteAddr = remoteAddr
}
req.Header.Set("X-Agent-Token", token)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected status %d for %s %s: %s", recorder.Code, method, path, recorder.Body.String())
}
var resp apiResponse
if err = json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if !resp.Success {
t.Fatalf("request %s %s failed: %s", method, path, resp.Message)
}
return resp
}
func createRouteAndPublishVersion(t *testing.T, engine http.Handler, adminToken string) {
t.Helper()
createBody := map[string]any{
"domain": "agent.example.com",
"origin_url": "https://agent-origin.internal",
"enabled": true,
"remark": "agent route",
}
performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/proxy-routes/", createBody)
performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/config-versions/publish", nil)
}

Some files were not shown because too many files have changed in this diff Show More