[优化] go 引用调整

This commit is contained in:
ryan
2026-06-06 10:26:20 +08:00
parent ee1110b752
commit 3cfefb4367
552 changed files with 1642 additions and 2185 deletions
+16
View File
@@ -0,0 +1,16 @@
.git
.github
node_modules
web/node_modules
web/.next
web/build
web/out
dist
tmp
logs
upload
*.log
.DS_Store
.env
.env.*
docker-compose*.yml
+1
View File
@@ -0,0 +1 @@
/data/
+38
View File
@@ -0,0 +1,38 @@
# syntax=docker/dockerfile:1.7
ARG VERSION=dev
FROM node:20 AS web-builder
ARG VERSION
WORKDIR /build
RUN corepack enable
COPY openflare-server/web/package.json openflare-server/web/pnpm-lock.yaml ./
RUN --mount=type=cache,id=pnpm-store,target=/root/.local/share/pnpm/store \
pnpm install --frozen-lockfile
COPY openflare-server/web ./
RUN --mount=type=cache,id=next-cache,target=/build/.next/cache \
NEXT_PUBLIC_APP_VERSION="$VERSION" pnpm build
FROM golang:1.25 AS go-builder
ARG VERSION
ENV GO111MODULE=on \
CGO_ENABLED=0 \
GOOS=linux
WORKDIR /build
COPY go.mod go.sum ./
RUN --mount=type=cache,target=/go/pkg/mod \
go mod download
COPY openflare-server ./openflare-server
COPY --from=web-builder /build/build ./openflare-server/web/build
RUN --mount=type=cache,target=/go/pkg/mod \
--mount=type=cache,target=/root/.cache/go-build \
go build -trimpath \
-ldflags "-s -w -X 'github.com/rain-kl/openflare/openflare-server/common.Version=$VERSION'" \
-o openflare ./openflare-server
FROM alpine:latest
RUN apk add --no-cache ca-certificates tzdata \
&& update-ca-certificates 2>/dev/null || true
ENV PORT=3000
COPY --from=go-builder /build/openflare /openflare
EXPOSE 3000
WORKDIR /data
ENTRYPOINT ["/openflare"]
+177
View File
@@ -0,0 +1,177 @@
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 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,
})
}
+196
View File
@@ -0,0 +1,196 @@
package controller
import (
"strconv"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/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/common/response"
"github.com/rain-kl/openflare/openflare-server/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)
}
+363
View File
@@ -0,0 +1,363 @@
package controller
import (
"encoding/json"
"log/slog"
"net"
"strconv"
"time"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/model"
"github.com/rain-kl/openflare/openflare-server/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
}
+358
View File
@@ -0,0 +1,358 @@
package controller
import (
"encoding/json"
"fmt"
"net/url"
"strconv"
"strings"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/model"
"github.com/rain-kl/openflare/openflare-server/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
}
+49
View File
@@ -0,0 +1,49 @@
package bind
import (
"encoding/json"
"errors"
"io"
"strconv"
"github.com/gin-gonic/gin"
"github.com/rain-kl/openflare/openflare-server/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,165 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/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": "清理成功",
})
}
+173
View File
@@ -0,0 +1,173 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/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
}
+31
View File
@@ -0,0 +1,31 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/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)
}
+135
View File
@@ -0,0 +1,135 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/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)
}
+176
View File
@@ -0,0 +1,176 @@
package controller
import (
"log/slog"
"net"
"time"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/model"
"github.com/rain-kl/openflare/openflare-server/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
}
+37
View File
@@ -0,0 +1,37 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/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)
}
+155
View File
@@ -0,0 +1,155 @@
package controller
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"log/slog"
"net/http"
"time"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/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/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/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)
}
+147
View File
@@ -0,0 +1,147 @@
package controller
import (
"fmt"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/model"
"github.com/rain-kl/openflare/openflare-server/service"
"github.com/rain-kl/openflare/openflare-server/utils/mail"
"github.com/rain-kl/openflare/openflare-server/utils/security"
"github.com/rain-kl/openflare/openflare-server/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,
"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)
}
+288
View File
@@ -0,0 +1,288 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/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)
}
+417
View File
@@ -0,0 +1,417 @@
package controller
import (
"fmt"
"regexp"
"strconv"
"strings"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/model"
"github.com/rain-kl/openflare/openflare-server/service"
"github.com/rain-kl/openflare/openflare-server/utils"
"github.com/rain-kl/openflare/openflare-server/utils/geoip"
"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, "")
}
+119
View File
@@ -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")
}
}
+73
View File
@@ -0,0 +1,73 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/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)
}
+170
View File
@@ -0,0 +1,170 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/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)
}
+119
View File
@@ -0,0 +1,119 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/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)
}
+121
View File
@@ -0,0 +1,121 @@
package controller
import (
"log/slog"
"net"
"time"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/model"
"github.com/rain-kl/openflare/openflare-server/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/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/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)
}
+166
View File
@@ -0,0 +1,166 @@
package controller
import (
"time"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/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": "服务升级任务已启动,确认无误后将自动重启。",
//})
}
+25
View File
@@ -0,0 +1,25 @@
package controller
import (
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/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, "同步成功")
}
+415
View File
@@ -0,0 +1,415 @@
package controller
import (
"strconv"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/middleware"
"github.com/rain-kl/openflare/openflare-server/model"
"github.com/rain-kl/openflare/openflare-server/utils/security"
"github.com/rain-kl/openflare/openflare-server/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/common/response"
"github.com/rain-kl/openflare/openflare-server/controller/bind"
"github.com/rain-kl/openflare/openflare-server/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
}
+125
View File
@@ -0,0 +1,125 @@
package controller
import (
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"time"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/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
}
+35
View File
@@ -0,0 +1,35 @@
services:
postgres:
image: postgres:17-alpine
restart: unless-stopped
environment:
POSTGRES_DB: openflare
POSTGRES_USER: openflare
POSTGRES_PASSWORD: replace-with-strong-password
volumes:
- ./postgres-data:/var/lib/postgresql/data
healthcheck:
test: ["CMD-SHELL", "pg_isready -U openflare -d openflare"]
interval: 10s
timeout: 5s
retries: 5
openflare:
build:
dockerfile: Dockerfile
restart: unless-stopped
depends_on:
postgres:
condition: service_healthy
ports:
- "3000:3000"
environment:
SESSION_SECRET: replace-with-random-string
SQLITE_PATH: /data/openflare.db
DSN: postgres://openflare:replace-with-strong-password@postgres:5432/openflare?sslmode=disable
GIN_MODE: release
LOG_LEVEL: info
volumes:
- ./openflare-data:/data
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+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()
}
}
+44
View File
@@ -0,0 +1,44 @@
package job
import (
"log/slog"
"time"
"github.com/rain-kl/openflare/openflare-server/model"
"github.com/rain-kl/openflare/openflare-server/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")
}
+44
View File
@@ -0,0 +1,44 @@
package job
import (
"log/slog"
"sync"
"time"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/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")
}
}
+15
View File
@@ -0,0 +1,15 @@
package job
import (
"log/slog"
"github.com/rain-kl/openflare/openflare-server/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)
}
}
+123
View File
@@ -0,0 +1,123 @@
package main
import (
"context"
"embed"
"fmt"
"log/slog"
"os"
"strconv"
"github.com/rain-kl/openflare/openflare-server/common"
_ "github.com/rain-kl/openflare/openflare-server/docs"
"github.com/rain-kl/openflare/openflare-server/job"
"github.com/rain-kl/openflare/openflare-server/middleware"
"github.com/rain-kl/openflare/openflare-server/model"
"github.com/rain-kl/openflare/openflare-server/router"
"github.com/rain-kl/openflare/openflare-server/service"
"github.com/rain-kl/openflare/openflare-server/utils/geoip"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-contrib/sessions/redis"
"github.com/gin-gonic/gin"
)
//go:embed all:web/build
var buildFS embed.FS
//go:embed web/build/index.html
var indexPage []byte
// @title OpenFlare Server API
// @version 3.0
// @description OpenFlare Server 管理端与 Agent API 文档。
// @BasePath /
// @schemes http https
// @securityDefinitions.apikey OpenFlareTokenAuth
// @in header
// @name OpenFlare-Token
// @description 管理端 API 使用登录后返回的用户 Token
// @securityDefinitions.apikey AccessTokenAuth
// @in header
// @name X-Agent-Token
// @description Agent API 使用节点专属 Agent Token 或全局 Discovery Token
func main() {
common.ParseFlags()
common.SetupGinLog()
slog.Info("OpenFlare started", "version", common.Version)
if os.Getenv("GIN_MODE") != "debug" {
gin.SetMode(gin.ReleaseMode)
}
// Initialize SQL Database
defer service.ShutdownWSHubs()
err := model.InitDB()
if err != nil {
slog.Error("initialize database failed", "error", err)
os.Exit(1)
}
defer func() {
err := model.CloseDB()
if err != nil {
slog.Error("close database failed", "error", err)
os.Exit(1)
}
}()
// Initialize Redis
err = common.InitRedisClient()
if err != nil {
slog.Error("initialize redis failed", "error", err)
os.Exit(1)
}
// Initialize options
model.InitOptionMap()
middleware.InitJWTMiddleware()
geoip.InitGeoIP(common.GeoIPProvider)
backgroundCtx, cancelBackgroundTasks := context.WithCancel(context.Background())
defer cancelBackgroundTasks()
service.StartDatabaseAutoCleanupScheduler(backgroundCtx)
job.InitCronJobs()
defer job.StopCronJobs()
// Initialize HTTP server
server := gin.Default()
//server.Use(gzip.Gzip(gzip.DefaultCompression))
server.Use(middleware.CORS())
// Initialize session store
if common.RedisEnabled {
opt := common.ParseRedisOption()
store, _ := redis.NewStore(opt.MinIdleConns, opt.Network, opt.Addr, opt.Password, []byte(common.SessionSecret))
server.Use(sessions.Sessions("session", store))
} else {
store := cookie.NewStore([]byte(common.SessionSecret))
server.Use(sessions.Sessions("session", store))
}
router.SetRouter(server, buildFS, indexPage)
var port = os.Getenv("PORT")
if port == "" {
port = strconv.Itoa(*common.Port)
}
dbBackend := "sqlite"
if common.SQLDSN != "" {
dbBackend = "postgres"
}
slog.Info("server config", "port", port, "gin_mode", gin.Mode(), "log_level", common.GetLogLevel(), "db_backend", dbBackend, "sqlite_path", common.SQLitePath, "redis_enabled", common.RedisEnabled, "log_dir", valueOrDefault(*common.LogDir, "stdout"), "access_token_configured", common.AccessToken != "", "node_offline_threshold", common.NodeOfflineThreshold)
slog.Info("server listening", "address", fmt.Sprintf(":%s", port))
err = server.Run(":" + port)
if err != nil {
slog.Error("server run failed", "error", err)
}
}
func valueOrDefault(value string, fallback string) string {
if value == "" {
return fallback
}
return value
}
+39
View File
@@ -0,0 +1,39 @@
package middleware
import (
"github.com/gin-gonic/gin"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/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()
}
}
+101
View File
@@ -0,0 +1,101 @@
package middleware
import (
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/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()
}
}
+37
View File
@@ -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
}
}
+27
View File
@@ -0,0 +1,27 @@
package middleware
import (
"strings"
"github.com/rain-kl/openflare/openflare-server/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)
}
+72
View File
@@ -0,0 +1,72 @@
package middleware
import (
"log"
"time"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/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())
}
}
+98
View File
@@ -0,0 +1,98 @@
package middleware
import (
"context"
"log/slog"
"net/http"
"time"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/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")
}
+28
View File
@@ -0,0 +1,28 @@
package middleware
import (
"github.com/gin-gonic/gin"
"github.com/rain-kl/openflare/openflare-server/common/response"
"github.com/rain-kl/openflare/openflare-server/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/common/response"
"github.com/rain-kl/openflare/openflare-server/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()
}
}
+41
View File
@@ -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
}
+81
View File
@@ -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
}
+279
View File
@@ -0,0 +1,279 @@
package model
import (
"errors"
"regexp"
"strings"
"time"
"github.com/rain-kl/openflare/openflare-server/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
}
+45
View File
@@ -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/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"
}
+35
View File
@@ -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
}
+244
View File
@@ -0,0 +1,244 @@
package goose
import (
"database/sql"
"errors"
"fmt"
"log/slog"
"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) error {
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
}
@@ -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
}
+81
View File
@@ -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/common"
"github.com/rain-kl/openflare/openflare-server/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")
}
+886
View File
@@ -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")
}
}
+41
View File
@@ -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)
}
}
+26
View File
@@ -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)
}
+26
View File
@@ -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)
}
+26
View File
@@ -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)
}
+26
View File
@@ -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)
}
+26
View File
@@ -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)
}
+52
View File
@@ -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
}
+185
View File
@@ -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
}
+58
View File
@@ -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
}
+41
View File
@@ -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)
}
+29
View File
@@ -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
}
+749
View File
@@ -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/openflare-server/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/openflare-server/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/openflare-server/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/openflare-server/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/openflare-server/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
}
+423
View File
@@ -0,0 +1,423 @@
package model
import (
"strconv"
"strings"
"time"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/utils/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["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 "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
}
+101
View File
@@ -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/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
}
+51
View File
@@ -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/common"
"github.com/rain-kl/openflare/openflare-server/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
}

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