mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-30 22:26:38 +08:00
[优化] go 引用调整
This commit is contained in:
@@ -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
|
||||
@@ -0,0 +1 @@
|
||||
/data/
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
)
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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, " ")
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
"github.com/go-redis/redis/v8"
|
||||
)
|
||||
|
||||
var RDB *redis.Client
|
||||
var RedisEnabled = true
|
||||
|
||||
// InitRedisClient This function is called after init()
|
||||
func InitRedisClient() (err error) {
|
||||
if os.Getenv("REDIS_CONN_STRING") == "" {
|
||||
RedisEnabled = false
|
||||
slog.Info("redis disabled because REDIS_CONN_STRING is not set")
|
||||
return nil
|
||||
}
|
||||
opt, err := redis.ParseURL(os.Getenv("REDIS_CONN_STRING"))
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
RDB = redis.NewClient(opt)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
_, err = RDB.Ping(ctx).Result()
|
||||
return err
|
||||
}
|
||||
|
||||
func ParseRedisOption() *redis.Options {
|
||||
opt, err := redis.ParseURL(os.Getenv("REDIS_CONN_STRING"))
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return opt
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
package response
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
const invalidParamsMessage = "参数错误"
|
||||
|
||||
// RespondSuccess sends a successful response with data
|
||||
func RespondSuccess(c *gin.Context, data any) {
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"message": "",
|
||||
"data": data,
|
||||
})
|
||||
}
|
||||
|
||||
// RespondSuccessWithExtras sends a successful response with data and extra fields
|
||||
func RespondSuccessWithExtras(c *gin.Context, data any, extras gin.H) {
|
||||
payload := gin.H{
|
||||
"success": true,
|
||||
"message": "",
|
||||
"data": data,
|
||||
}
|
||||
for key, value := range extras {
|
||||
payload[key] = value
|
||||
}
|
||||
c.JSON(http.StatusOK, payload)
|
||||
}
|
||||
|
||||
// RespondSuccessMessage sends a successful response with a custom message
|
||||
func RespondSuccessMessage(c *gin.Context, message string) {
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"message": message,
|
||||
})
|
||||
}
|
||||
|
||||
// RespondFailure sends a failed response with http.StatusOK and a failure message
|
||||
func RespondFailure(c *gin.Context, message string) {
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": false,
|
||||
"message": message,
|
||||
})
|
||||
}
|
||||
|
||||
// RespondBadRequest sends a bad request response (400)
|
||||
func RespondBadRequest(c *gin.Context, message string) {
|
||||
if message == "" {
|
||||
message = invalidParamsMessage
|
||||
}
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"success": false,
|
||||
"message": message,
|
||||
})
|
||||
}
|
||||
|
||||
// RespondUnauthorized sends an unauthorized response (401)
|
||||
func RespondUnauthorized(c *gin.Context, message string) {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{
|
||||
"success": false,
|
||||
"message": message,
|
||||
})
|
||||
}
|
||||
|
||||
// RespondForbidden sends a forbidden response (403)
|
||||
func RespondForbidden(c *gin.Context, message string) {
|
||||
c.JSON(http.StatusForbidden, gin.H{
|
||||
"success": false,
|
||||
"message": message,
|
||||
})
|
||||
}
|
||||
|
||||
// RespondErrorWithStatus sends a response with target HTTP status code and a message
|
||||
func RespondErrorWithStatus(c *gin.Context, code int, message string) {
|
||||
c.JSON(code, gin.H{
|
||||
"success": false,
|
||||
"message": message,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 = ¤tUser.Id
|
||||
}
|
||||
result, pending, err := service.CompleteOAuthLogin(source, profile, currentUserID)
|
||||
if err != nil {
|
||||
response.RespondFailure(c, err.Error())
|
||||
return
|
||||
}
|
||||
if pending != nil {
|
||||
raw, err := json.Marshal(pending)
|
||||
if err != nil {
|
||||
response.RespondFailure(c, err.Error())
|
||||
return
|
||||
}
|
||||
session.Set(pendingExternalAccountSessionKey, string(raw))
|
||||
if err := session.Save(); err != nil {
|
||||
response.RespondFailure(c, "无法保存待绑定账号,请重试")
|
||||
return
|
||||
}
|
||||
response.RespondSuccess(c, result)
|
||||
return
|
||||
}
|
||||
if result.User != nil {
|
||||
cleanUser, err := setLoginToken(result.User)
|
||||
if err != nil {
|
||||
response.RespondFailure(c, "无法保存会话信息,请重试")
|
||||
return
|
||||
}
|
||||
result.User = cleanUser
|
||||
}
|
||||
response.RespondSuccess(c, result)
|
||||
}
|
||||
|
||||
func LinkExistingOAuthAccount(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
raw, _ := session.Get(pendingExternalAccountSessionKey).(string)
|
||||
if raw == "" {
|
||||
response.RespondFailure(c, "待绑定第三方账号已失效,请重新登录")
|
||||
return
|
||||
}
|
||||
var pending service.PendingExternalAccount
|
||||
if err := json.Unmarshal([]byte(raw), &pending); err != nil {
|
||||
response.RespondFailure(c, "待绑定第三方账号无效,请重新登录")
|
||||
return
|
||||
}
|
||||
var input service.LinkExistingRequest
|
||||
if err := bind.DecodeJSONBody(c.Request.Body, &input); err != nil {
|
||||
response.RespondBadRequest(c, "无效的参数")
|
||||
return
|
||||
}
|
||||
user, err := service.LinkPendingExternalAccount(&pending, input)
|
||||
if err != nil {
|
||||
response.RespondFailure(c, err.Error())
|
||||
return
|
||||
}
|
||||
session.Delete(pendingExternalAccountSessionKey)
|
||||
if err := session.Save(); err != nil {
|
||||
response.RespondFailure(c, "无法更新会话信息,请重试")
|
||||
return
|
||||
}
|
||||
cleanUser, err := setLoginToken(user)
|
||||
if err != nil {
|
||||
response.RespondFailure(c, "无法保存会话信息,请重试")
|
||||
return
|
||||
}
|
||||
response.RespondSuccess(c, service.OAuthCallbackResult{Status: "linked", User: cleanUser})
|
||||
}
|
||||
|
||||
func ListExternalAccounts(c *gin.Context) {
|
||||
userID := c.GetInt("id")
|
||||
accounts, err := model.ListExternalAccountsByUserID(userID)
|
||||
if err != nil {
|
||||
response.RespondFailure(c, err.Error())
|
||||
return
|
||||
}
|
||||
response.RespondSuccess(c, accounts)
|
||||
}
|
||||
|
||||
func DeleteExternalAccount(c *gin.Context) {
|
||||
rawID := strings.TrimSpace(c.Param("id"))
|
||||
parsedID, err := strconv.ParseUint(rawID, 10, 64)
|
||||
if err != nil || parsedID == 0 {
|
||||
response.RespondBadRequest(c, "绑定记录 ID 无效")
|
||||
return
|
||||
}
|
||||
if err := model.DeleteExternalAccountForUser(uint(parsedID), c.GetInt("id")); err != nil {
|
||||
response.RespondFailure(c, err.Error())
|
||||
return
|
||||
}
|
||||
response.RespondSuccessMessage(c, "")
|
||||
}
|
||||
|
||||
func parseAuthSourceID(c *gin.Context) (uint, error) {
|
||||
raw := c.Param("source_id")
|
||||
if raw == "" {
|
||||
raw = c.Param("id")
|
||||
}
|
||||
parsed, err := strconv.ParseUint(raw, 10, 64)
|
||||
if err != nil || parsed == 0 {
|
||||
return 0, fmt.Errorf("认证源 ID 无效")
|
||||
}
|
||||
return uint(parsed), nil
|
||||
}
|
||||
|
||||
func getAuthSourceFromRoute(c *gin.Context) (*model.AuthSource, error) {
|
||||
raw := strings.TrimSpace(c.Param("source"))
|
||||
if raw == "" {
|
||||
raw = strings.TrimSpace(c.Param("source_id"))
|
||||
}
|
||||
if raw == "" {
|
||||
raw = strings.TrimSpace(c.Param("id"))
|
||||
}
|
||||
if raw == "" {
|
||||
return nil, fmt.Errorf("认证源不能为空")
|
||||
}
|
||||
if parsed, err := strconv.ParseUint(raw, 10, 64); err == nil && parsed > 0 {
|
||||
source, err := model.GetAuthSourceByID(uint(parsed))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return source, nil
|
||||
}
|
||||
source, err := model.GetAuthSourceByName(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return source, nil
|
||||
}
|
||||
|
||||
func oauthStateSessionKey(sourceID uint) string {
|
||||
return fmt.Sprintf("oauth_state_%d", sourceID)
|
||||
}
|
||||
|
||||
func oauthFrontendCallbackURL(c *gin.Context, sourceID uint) string {
|
||||
base := strings.TrimRight(common.ServerAddress, "/")
|
||||
if base == "" {
|
||||
scheme := "http"
|
||||
if c.Request.TLS != nil || c.GetHeader("X-Forwarded-Proto") == "https" {
|
||||
scheme = "https"
|
||||
}
|
||||
host := c.Request.Host
|
||||
if forwardedHost := c.GetHeader("X-Forwarded-Host"); forwardedHost != "" {
|
||||
host = forwardedHost
|
||||
}
|
||||
base = scheme + "://" + host
|
||||
}
|
||||
source, err := model.GetAuthSourceByID(sourceID)
|
||||
sourceName := strconv.FormatUint(uint64(sourceID), 10)
|
||||
if err == nil && strings.TrimSpace(source.Name) != "" {
|
||||
sourceName = source.Name
|
||||
}
|
||||
callback, _ := url.JoinPath(base, "oauth", sourceName)
|
||||
return callback
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
package bind
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"strconv"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/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": "清理成功",
|
||||
})
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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, "")
|
||||
}
|
||||
@@ -0,0 +1,119 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestValidateOpenRestyOption(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
key string
|
||||
value string
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "default server status valid 421", key: "OpenRestyDefaultServerReturnStatus", value: "421"},
|
||||
{name: "default server status valid 200", key: "OpenRestyDefaultServerReturnStatus", value: "200"},
|
||||
{name: "default server status invalid 99", key: "OpenRestyDefaultServerReturnStatus", value: "99", wantErr: true},
|
||||
{name: "default server status invalid 1000", key: "OpenRestyDefaultServerReturnStatus", value: "1000", wantErr: true},
|
||||
{name: "default server status invalid abc", key: "OpenRestyDefaultServerReturnStatus", value: "abc", wantErr: true},
|
||||
{name: "worker processes auto", key: "OpenRestyWorkerProcesses", value: "auto"},
|
||||
{name: "worker processes number", key: "OpenRestyWorkerProcesses", value: "8"},
|
||||
{name: "worker processes invalid", key: "OpenRestyWorkerProcesses", value: "0", wantErr: true},
|
||||
{name: "events use empty", key: "OpenRestyEventsUse", value: ""},
|
||||
{name: "events use invalid", key: "OpenRestyEventsUse", value: "io_uring", wantErr: true},
|
||||
{name: "resolvers valid", key: "OpenRestyResolvers", value: "1.1.1.1 8.8.8.8"},
|
||||
{name: "resolvers invalid", key: "OpenRestyResolvers", value: "1.1.1.1; 8.8.8.8", wantErr: true},
|
||||
{name: "proxy buffers valid", key: "OpenRestyProxyBuffers", value: "16 16k"},
|
||||
{name: "proxy buffers invalid", key: "OpenRestyProxyBuffers", value: "16x16k", wantErr: true},
|
||||
{name: "cache max size valid", key: "OpenRestyCacheMaxSize", value: "2g"},
|
||||
{name: "cache max size invalid", key: "OpenRestyCacheMaxSize", value: "2gb", wantErr: true},
|
||||
{name: "client max body size valid", key: "OpenRestyClientMaxBodySize", value: "64m"},
|
||||
{name: "client max body size invalid", key: "OpenRestyClientMaxBodySize", value: "64mb", wantErr: true},
|
||||
{name: "large client header buffers valid", key: "OpenRestyLargeClientHeaderBuffers", value: "4 16k"},
|
||||
{name: "large client header buffers invalid", key: "OpenRestyLargeClientHeaderBuffers", value: "4x16k", wantErr: true},
|
||||
{name: "proxy request buffering valid", key: "OpenRestyProxyRequestBufferingEnabled", value: "true"},
|
||||
{name: "proxy request buffering invalid", key: "OpenRestyProxyRequestBufferingEnabled", value: "on", wantErr: true},
|
||||
{name: "websocket valid", key: "OpenRestyWebsocketEnabled", value: "false"},
|
||||
{name: "websocket invalid", key: "OpenRestyWebsocketEnabled", value: "off", wantErr: true},
|
||||
{name: "cache inactive valid", key: "OpenRestyCacheInactive", value: "30m"},
|
||||
{name: "cache inactive invalid", key: "OpenRestyCacheInactive", value: "30", wantErr: true},
|
||||
{name: "cache use stale valid", key: "OpenRestyCacheUseStale", value: "error timeout http_500"},
|
||||
{name: "cache use stale invalid", key: "OpenRestyCacheUseStale", value: "error whatever", wantErr: true},
|
||||
{name: "gzip level valid", key: "OpenRestyGzipCompLevel", value: "9"},
|
||||
{name: "gzip level invalid", key: "OpenRestyGzipCompLevel", value: "10", wantErr: true},
|
||||
}
|
||||
|
||||
for _, testCase := range testCases {
|
||||
err := validateOpenRestyOption(testCase.key, testCase.value)
|
||||
if testCase.wantErr && err == nil {
|
||||
t.Fatalf("%s: expected error", testCase.name)
|
||||
}
|
||||
if !testCase.wantErr && err != nil {
|
||||
t.Fatalf("%s: unexpected error: %v", testCase.name, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateAgentOption(t *testing.T) {
|
||||
if err := validateAgentOption("AgentWebsocketUpgradeEnabled", "true"); err != nil {
|
||||
t.Fatalf("expected websocket upgrade option to accept true: %v", err)
|
||||
}
|
||||
if err := validateAgentOption("AgentWebsocketUpgradeEnabled", "false"); err != nil {
|
||||
t.Fatalf("expected websocket upgrade option to accept false: %v", err)
|
||||
}
|
||||
if err := validateAgentOption("AgentWebsocketUpgradeEnabled", "on"); err == nil {
|
||||
t.Fatal("expected websocket upgrade option to reject non-boolean value")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateUptimeKumaOption(t *testing.T) {
|
||||
state := map[string]string{
|
||||
"UptimeKumaUrl": "http://localhost:3001",
|
||||
"UptimeKumaUsername": "admin",
|
||||
"UptimeKumaPassword": "password",
|
||||
}
|
||||
|
||||
testCases := []struct {
|
||||
name string
|
||||
key string
|
||||
value string
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "enabled true", key: "UptimeKumaEnabled", value: "true"},
|
||||
{name: "enabled false", key: "UptimeKumaEnabled", value: "false"},
|
||||
{name: "enabled invalid", key: "UptimeKumaEnabled", value: "on", wantErr: true},
|
||||
{name: "url http valid", key: "UptimeKumaUrl", value: "http://192.168.1.100:3001"},
|
||||
{name: "url https valid", key: "UptimeKumaUrl", value: "https://kuma.example.com"},
|
||||
{name: "url invalid", key: "UptimeKumaUrl", value: "kuma.example.com", wantErr: true},
|
||||
{name: "scope all", key: "UptimeKumaMonitorScope", value: "all"},
|
||||
{name: "scope selected", key: "UptimeKumaMonitorScope", value: "selected"},
|
||||
{name: "scope invalid", key: "UptimeKumaMonitorScope", value: "none", wantErr: true},
|
||||
{name: "sync interval valid", key: "UptimeKumaSyncInterval", value: "5"},
|
||||
{name: "sync interval invalid", key: "UptimeKumaSyncInterval", value: "0", wantErr: true},
|
||||
{name: "interval valid", key: "UptimeKumaInterval", value: "60"},
|
||||
{name: "interval invalid", key: "UptimeKumaInterval", value: "-60", wantErr: true},
|
||||
{name: "retry valid", key: "UptimeKumaRetry", value: "0"},
|
||||
{name: "retry positive valid", key: "UptimeKumaRetry", value: "3"},
|
||||
{name: "retry invalid", key: "UptimeKumaRetry", value: "-1", wantErr: true},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
err := validateUptimeKumaOption(tc.key, tc.value, state)
|
||||
if tc.wantErr && err == nil {
|
||||
t.Fatalf("%s: expected error", tc.name)
|
||||
}
|
||||
if !tc.wantErr && err != nil {
|
||||
t.Fatalf("%s: unexpected error: %v", tc.name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Test enabling Uptime Kuma when URL or credentials are empty in state
|
||||
stateEmpty := map[string]string{
|
||||
"UptimeKumaUrl": "",
|
||||
"UptimeKumaUsername": "",
|
||||
"UptimeKumaPassword": "",
|
||||
}
|
||||
if err := validateUptimeKumaOption("UptimeKumaEnabled", "true", stateEmpty); err == nil {
|
||||
t.Fatal("expected error when enabling Uptime Kuma with empty URL/credentials in state")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"github.com/rain-kl/openflare/openflare-server/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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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": "服务升级任务已启动,确认无误后将自动重启。",
|
||||
//})
|
||||
}
|
||||
@@ -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, "同步成功")
|
||||
}
|
||||
@@ -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, "")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -0,0 +1,44 @@
|
||||
package job
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
|
||||
"github.com/robfig/cron/v3"
|
||||
)
|
||||
|
||||
var cronRunner *cron.Cron
|
||||
|
||||
func InitCronJobs() {
|
||||
cronRunner = cron.New()
|
||||
|
||||
// Register SSL renew job
|
||||
_, err := cronRunner.AddJob("0 0 * * *", &SSLRenewJob{})
|
||||
if err != nil {
|
||||
slog.Error("failed to register SSL renew cron job", "error", err)
|
||||
} else {
|
||||
slog.Info("registered SSL renew cron job")
|
||||
}
|
||||
|
||||
_, err = cronRunner.AddJob("@every 5m", &WAFIPGroupSyncJob{})
|
||||
if err != nil {
|
||||
slog.Error("failed to register WAF IP group sync cron job", "error", err)
|
||||
} else {
|
||||
slog.Info("registered WAF IP group sync cron job")
|
||||
}
|
||||
|
||||
// Register Uptime Kuma sync job (check every minute)
|
||||
_, err = cronRunner.AddJob("* * * * *", &UptimeKumaSyncJob{})
|
||||
if err != nil {
|
||||
slog.Error("failed to register Uptime Kuma sync cron job", "error", err)
|
||||
} else {
|
||||
slog.Info("registered Uptime Kuma sync cron job")
|
||||
}
|
||||
|
||||
cronRunner.Start()
|
||||
}
|
||||
|
||||
func StopCronJobs() {
|
||||
if cronRunner != nil {
|
||||
cronRunner.Stop()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
package job
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/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")
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"path"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func Cache() func(c *gin.Context) {
|
||||
return func(c *gin.Context) {
|
||||
requestPath := c.Request.URL.Path
|
||||
|
||||
switch {
|
||||
case strings.HasPrefix(requestPath, "/_next/static/"):
|
||||
c.Header("Cache-Control", "public, max-age=31536000, immutable")
|
||||
case isStaticPublicAsset(requestPath):
|
||||
c.Header("Cache-Control", "public, max-age=86400")
|
||||
default:
|
||||
c.Header("Cache-Control", "no-store, no-cache, must-revalidate")
|
||||
c.Header("Pragma", "no-cache")
|
||||
c.Header("Expires", "0")
|
||||
}
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
func isStaticPublicAsset(requestPath string) bool {
|
||||
ext := strings.ToLower(path.Ext(requestPath))
|
||||
switch ext {
|
||||
case ".ico", ".png", ".jpg", ".jpeg", ".gif", ".svg", ".webp", ".css", ".js":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,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)
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type AcmeAccount struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
Email string `json:"email" gorm:"size:255"`
|
||||
URL string `json:"url" gorm:"size:255"`
|
||||
PrivateKey string `json:"-" gorm:"type:text;not null"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
func GetAcmeAccountByID(id uint) (*AcmeAccount, error) {
|
||||
account := &AcmeAccount{}
|
||||
err := DB.First(account, id).Error
|
||||
return account, err
|
||||
}
|
||||
|
||||
func GetDefaultAcmeAccount() (*AcmeAccount, error) {
|
||||
account := &AcmeAccount{}
|
||||
err := DB.Order("id asc").First(account).Error
|
||||
if err != nil {
|
||||
// Auto-create a default account placeholder if none exists
|
||||
account.Email = "admin@openflare.dev"
|
||||
err = DB.Create(account).Error
|
||||
}
|
||||
return account, err
|
||||
}
|
||||
|
||||
func (account *AcmeAccount) Insert() error {
|
||||
return DB.Create(account).Error
|
||||
}
|
||||
|
||||
func (account *AcmeAccount) Update() error {
|
||||
return DB.Save(account).Error
|
||||
}
|
||||
|
||||
func (account *AcmeAccount) Delete() error {
|
||||
return DB.Delete(account).Error
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type ApplyLogQuery struct {
|
||||
NodeID string
|
||||
PageNo int
|
||||
PageSize int
|
||||
}
|
||||
|
||||
type ApplyLog struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
|
||||
Version string `json:"version" gorm:"size:32;not null"`
|
||||
Result string `json:"result" gorm:"size:32;not null"`
|
||||
Message string `json:"message" gorm:"type:text"`
|
||||
Checksum string `json:"checksum" gorm:"size:64;not null;default:''"`
|
||||
MainConfigChecksum string `json:"main_config_checksum" gorm:"size:64;not null;default:''"`
|
||||
RouteConfigChecksum string `json:"route_config_checksum" gorm:"size:64;not null;default:''"`
|
||||
SupportFileCount int `json:"support_file_count" gorm:"not null;default:0"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
func ListApplyLogs(query ApplyLogQuery) (logs []*ApplyLog, err error) {
|
||||
db := DB.Order("id desc")
|
||||
if query.NodeID != "" {
|
||||
db = db.Where("node_id = ?", query.NodeID)
|
||||
}
|
||||
if query.PageSize > 0 {
|
||||
offset := 0
|
||||
if query.PageNo > 1 {
|
||||
offset = (query.PageNo - 1) * query.PageSize
|
||||
}
|
||||
db = db.Limit(query.PageSize).Offset(offset)
|
||||
}
|
||||
err = db.Find(&logs).Error
|
||||
return logs, err
|
||||
}
|
||||
|
||||
func CountApplyLogs(nodeID string) (total int64, err error) {
|
||||
query := DB.Model(&ApplyLog{})
|
||||
if nodeID != "" {
|
||||
query = query.Where("node_id = ?", nodeID)
|
||||
}
|
||||
err = query.Count(&total).Error
|
||||
return total, err
|
||||
}
|
||||
|
||||
func GetLatestApplyLogsByNodeIDs(nodeIDs []string) (map[string]*ApplyLog, error) {
|
||||
result := make(map[string]*ApplyLog)
|
||||
if len(nodeIDs) == 0 {
|
||||
return result, nil
|
||||
}
|
||||
|
||||
var logs []*ApplyLog
|
||||
subQuery := DB.Model(&ApplyLog{}).
|
||||
Select("MAX(id) AS id").
|
||||
Where("node_id IN ?", nodeIDs).
|
||||
Group("node_id")
|
||||
if err := DB.Where("id IN (?)", subQuery).Find(&logs).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, log := range logs {
|
||||
result[log.NodeID] = log
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func DeleteAllApplyLogs() (deleted int64, err error) {
|
||||
result := DB.Session(&gorm.Session{AllowGlobalUpdate: true}).Delete(&ApplyLog{})
|
||||
return result.RowsAffected, result.Error
|
||||
}
|
||||
|
||||
func DeleteApplyLogsBefore(before time.Time) (deleted int64, err error) {
|
||||
result := DB.Where("created_at < ?", before).Delete(&ApplyLog{})
|
||||
return result.RowsAffected, result.Error
|
||||
}
|
||||
@@ -0,0 +1,279 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/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(¤t, "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(¤t).Updates(map[string]any{
|
||||
"name": source.Name,
|
||||
"type": source.Type,
|
||||
"display_name": source.DisplayName,
|
||||
"is_active": source.IsActive,
|
||||
"client_id": source.ClientID,
|
||||
"client_secret": source.ClientSecret,
|
||||
"openid_discovery_url": source.OpenIDDiscoveryURL,
|
||||
"scopes": source.Scopes,
|
||||
"icon_url": source.IconURL,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func ToggleAuthSource(id uint, isActive bool) error {
|
||||
source, err := GetAuthSourceByID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
source.IsActive = isActive
|
||||
if err := source.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
return DB.Model(&AuthSource{}).Where("id = ?", id).Update("is_active", isActive).Error
|
||||
}
|
||||
|
||||
func DeleteAuthSource(id uint) error {
|
||||
if id == 0 {
|
||||
return errors.New("认证源 ID 不能为空")
|
||||
}
|
||||
return DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("auth_source_id = ?", id).Delete(&ExternalAccount{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(&AuthSource{}, "id = ?", id).Error
|
||||
})
|
||||
}
|
||||
|
||||
func FindExternalAccount(sourceID uint, externalID string) (*ExternalAccount, error) {
|
||||
var account ExternalAccount
|
||||
err := DB.Where("auth_source_id = ? AND external_id = ?", sourceID, externalID).First(&account).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &account, nil
|
||||
}
|
||||
|
||||
func LinkExternalAccount(account *ExternalAccount) error {
|
||||
if account.AuthSourceID == 0 || account.UserID == 0 || strings.TrimSpace(account.ExternalID) == "" {
|
||||
return errors.New("外部账号绑定信息不完整")
|
||||
}
|
||||
account.ExternalID = strings.TrimSpace(account.ExternalID)
|
||||
account.ExternalUsername = strings.TrimSpace(account.ExternalUsername)
|
||||
account.Email = strings.TrimSpace(account.Email)
|
||||
return DB.Where(ExternalAccount{
|
||||
AuthSourceID: account.AuthSourceID,
|
||||
ExternalID: account.ExternalID,
|
||||
}).FirstOrCreate(account).Error
|
||||
}
|
||||
|
||||
func ListExternalAccountsByUserID(userID int) ([]ExternalAccountView, error) {
|
||||
if userID <= 0 {
|
||||
return nil, errors.New("用户 ID 不能为空")
|
||||
}
|
||||
var accounts []ExternalAccount
|
||||
if err := DB.Preload("AuthSource").Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]ExternalAccountView, 0, len(accounts))
|
||||
for _, account := range accounts {
|
||||
label := account.AuthSource.DisplayName
|
||||
if label == "" {
|
||||
label = account.AuthSource.Name
|
||||
}
|
||||
views = append(views, ExternalAccountView{
|
||||
ID: account.ID,
|
||||
AuthSourceID: account.AuthSourceID,
|
||||
AuthSourceName: account.AuthSource.Name,
|
||||
AuthSourceType: account.AuthSource.Type,
|
||||
AuthSourceLabel: label,
|
||||
ExternalUsername: account.ExternalUsername,
|
||||
Email: account.Email,
|
||||
CreatedAt: account.CreatedAt,
|
||||
})
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
func DeleteExternalAccountForUser(id uint, userID int) error {
|
||||
if id == 0 {
|
||||
return errors.New("绑定记录 ID 不能为空")
|
||||
}
|
||||
if userID <= 0 {
|
||||
return errors.New("用户 ID 不能为空")
|
||||
}
|
||||
result := DB.Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{})
|
||||
if result.Error != nil {
|
||||
return result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return errors.New("绑定记录不存在")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type ConfigVersionSummary struct {
|
||||
ID uint `json:"id"`
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
IsActive bool `json:"is_active"`
|
||||
CreatedBy string `json:"created_by"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
type ConfigVersion struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
Version string `json:"version" gorm:"uniqueIndex;size:32;not null"`
|
||||
SnapshotJSON string `json:"snapshot_json" gorm:"type:text;not null"`
|
||||
MainConfig string `json:"main_config" gorm:"type:text;not null;default:''"`
|
||||
RenderedConfig string `json:"rendered_config" gorm:"type:text;not null"`
|
||||
SupportFilesJSON string `json:"support_files_json" gorm:"type:text;not null;default:'[]'"`
|
||||
Checksum string `json:"checksum" gorm:"size:64;not null"`
|
||||
IsActive bool `json:"is_active" gorm:"not null;default:false;index"`
|
||||
CreatedBy string `json:"created_by" gorm:"size:64;not null"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
func ListConfigVersionSummaries() (versions []*ConfigVersionSummary, err error) {
|
||||
err = DB.Model(&ConfigVersion{}).
|
||||
Select("id", "version", "checksum", "is_active", "created_by", "created_at").
|
||||
Order("id desc").
|
||||
Find(&versions).Error
|
||||
return versions, err
|
||||
}
|
||||
|
||||
func GetConfigVersionByID(id uint) (*ConfigVersion, error) {
|
||||
version := &ConfigVersion{}
|
||||
err := DB.First(version, id).Error
|
||||
return version, err
|
||||
}
|
||||
|
||||
func GetActiveConfigVersion() (*ConfigVersion, error) {
|
||||
version := &ConfigVersion{}
|
||||
err := DB.Where("is_active = ?", true).Order("id desc").First(version).Error
|
||||
return version, err
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model/migrate"
|
||||
)
|
||||
|
||||
const (
|
||||
legacyDatabaseSchemaVersion = migrate.BaseDatabaseSchemaVersion
|
||||
legacyMigrationTerminalVersion = 17
|
||||
databaseSchemaVersionRowID = 1
|
||||
)
|
||||
|
||||
// currentDatabaseSchemaVersion tracks the current physical schema validated by the
|
||||
// legacy validator set. Goose owns only post-v17 migrations, and none exist yet.
|
||||
var currentDatabaseSchemaVersion = legacyMigrationTerminalVersion
|
||||
|
||||
type DatabaseSchemaVersion struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
Version int `json:"version" gorm:"not null"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (DatabaseSchemaVersion) TableName() string {
|
||||
return "database_schema_versions"
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type DnsAccount struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
Name string `json:"name" gorm:"size:255;not null"`
|
||||
Type string `json:"type" gorm:"size:64;not null"`
|
||||
Authorization string `json:"-" gorm:"type:text;not null"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
func ListDnsAccounts() (accounts []*DnsAccount, err error) {
|
||||
err = DB.Order("id desc").Find(&accounts).Error
|
||||
return accounts, err
|
||||
}
|
||||
|
||||
func GetDnsAccountByID(id uint) (*DnsAccount, error) {
|
||||
account := &DnsAccount{}
|
||||
err := DB.First(account, id).Error
|
||||
return account, err
|
||||
}
|
||||
|
||||
func (account *DnsAccount) Insert() error {
|
||||
return DB.Create(account).Error
|
||||
}
|
||||
|
||||
func (account *DnsAccount) Update() error {
|
||||
return DB.Save(account).Error
|
||||
}
|
||||
|
||||
func (account *DnsAccount) Delete() error {
|
||||
return DB.Delete(account).Error
|
||||
}
|
||||
@@ -0,0 +1,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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -0,0 +1,886 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"go/ast"
|
||||
"go/parser"
|
||||
"go/token"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type legacyProxyRouteV7 struct {
|
||||
ID uint `gorm:"primaryKey"`
|
||||
SiteName string `gorm:"size:255;not null;default:''"`
|
||||
Domain string `gorm:"uniqueIndex;size:255;not null"`
|
||||
Domains string `gorm:"type:text;not null;default:'[]'"`
|
||||
OriginID *uint `gorm:"index"`
|
||||
OriginURL string `gorm:"size:2048;not null"`
|
||||
OriginHost string `gorm:"size:255"`
|
||||
Upstreams string `gorm:"type:text;not null;default:'[]'"`
|
||||
Enabled bool `gorm:"not null;default:true"`
|
||||
EnableHTTPS bool `gorm:"column:enable_https;not null;default:false"`
|
||||
CertID *uint
|
||||
CertIDs string `gorm:"type:text;not null;default:'[]'"`
|
||||
RedirectHTTP bool `gorm:"not null;default:false"`
|
||||
LimitConnPerServer int `gorm:"not null;default:0"`
|
||||
LimitConnPerIP int `gorm:"not null;default:0"`
|
||||
LimitRate string `gorm:"size:32;not null;default:''"`
|
||||
CacheEnabled bool `gorm:"not null;default:false"`
|
||||
CachePolicy string `gorm:"size:32;not null;default:''"`
|
||||
CacheRules string `gorm:"type:text;not null;default:'[]'"`
|
||||
CustomHeaders string `gorm:"type:text;not null;default:'[]'"`
|
||||
Remark string `gorm:"size:255"`
|
||||
CreatedAt time.Time
|
||||
UpdatedAt time.Time
|
||||
}
|
||||
|
||||
func (legacyProxyRouteV7) TableName() string {
|
||||
return "proxy_routes"
|
||||
}
|
||||
|
||||
func openBareTestSQLiteDB(t *testing.T, name string) *gorm.DB {
|
||||
t.Helper()
|
||||
|
||||
db, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), name)), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite db: %v", err)
|
||||
}
|
||||
sqlDB, err := db.DB()
|
||||
if err != nil {
|
||||
t.Fatalf("get sql db: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = sqlDB.Close()
|
||||
})
|
||||
return db
|
||||
}
|
||||
|
||||
func openTestSQLiteDB(t *testing.T, name string) *gorm.DB {
|
||||
t.Helper()
|
||||
|
||||
db := openBareTestSQLiteDB(t, name)
|
||||
if err := autoMigrateAll(db); err != nil {
|
||||
t.Fatalf("auto migrate db: %v", err)
|
||||
}
|
||||
return db
|
||||
}
|
||||
|
||||
func findDBModelByTableName(t *testing.T, tableName string) dbModel {
|
||||
t.Helper()
|
||||
|
||||
models, err := buildDBModels()
|
||||
if err != nil {
|
||||
t.Fatalf("build db models: %v", err)
|
||||
}
|
||||
for _, item := range models {
|
||||
if item.tableName == tableName {
|
||||
return item
|
||||
}
|
||||
}
|
||||
t.Fatalf("db model not found for table %s", tableName)
|
||||
return dbModel{}
|
||||
}
|
||||
|
||||
func expectedCurrentDatabaseVersion() int {
|
||||
return int(currentGooseTargetVersion())
|
||||
}
|
||||
|
||||
func TestIsDatabaseEmpty(t *testing.T) {
|
||||
db := openTestSQLiteDB(t, "empty.db")
|
||||
|
||||
empty, err := isDatabaseEmpty(db)
|
||||
if err != nil {
|
||||
t.Fatalf("isDatabaseEmpty returned error: %v", err)
|
||||
}
|
||||
if !empty {
|
||||
t.Fatal("expected database to be empty")
|
||||
}
|
||||
|
||||
if err := db.Create(&User{
|
||||
Username: "alice",
|
||||
Password: "secret",
|
||||
DisplayName: "Alice",
|
||||
Role: 1,
|
||||
Status: 1,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed user: %v", err)
|
||||
}
|
||||
|
||||
empty, err = isDatabaseEmpty(db)
|
||||
if err != nil {
|
||||
t.Fatalf("isDatabaseEmpty after seed returned error: %v", err)
|
||||
}
|
||||
if empty {
|
||||
t.Fatal("expected database to be non-empty")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrateTableDataCopiesRows(t *testing.T) {
|
||||
source := openTestSQLiteDB(t, "source.db")
|
||||
target := openTestSQLiteDB(t, "target.db")
|
||||
|
||||
user := User{
|
||||
Id: 1,
|
||||
Username: "root",
|
||||
Password: "hashed",
|
||||
DisplayName: "Root User",
|
||||
Role: 100,
|
||||
Status: 1,
|
||||
}
|
||||
option := Option{
|
||||
Key: "AgentHeartbeatInterval",
|
||||
Value: "10000",
|
||||
}
|
||||
|
||||
if err := source.Create(&user).Error; err != nil {
|
||||
t.Fatalf("seed source user: %v", err)
|
||||
}
|
||||
if err := source.Create(&option).Error; err != nil {
|
||||
t.Fatalf("seed source option: %v", err)
|
||||
}
|
||||
|
||||
if err := migrateTableData(source, target, findDBModelByTableName(t, "users")); err != nil {
|
||||
t.Fatalf("migrate users: %v", err)
|
||||
}
|
||||
if err := migrateTableData(source, target, findDBModelByTableName(t, "options")); err != nil {
|
||||
t.Fatalf("migrate options: %v", err)
|
||||
}
|
||||
|
||||
var gotUser User
|
||||
if err := target.First(&gotUser, 1).Error; err != nil {
|
||||
t.Fatalf("query migrated user: %v", err)
|
||||
}
|
||||
if gotUser.Username != user.Username || gotUser.DisplayName != user.DisplayName {
|
||||
t.Fatalf("unexpected migrated user: %+v", gotUser)
|
||||
}
|
||||
|
||||
var gotOption Option
|
||||
if err := target.First(&gotOption, "key = ?", option.Key).Error; err != nil {
|
||||
t.Fatalf("query migrated option: %v", err)
|
||||
}
|
||||
if gotOption.Value != option.Value {
|
||||
t.Fatalf("unexpected migrated option value: %s", gotOption.Value)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterShardingAutoMigratesShardTables(t *testing.T) {
|
||||
db := openBareTestSQLiteDB(t, "sharded.db")
|
||||
if err := registerSharding(db, "sqlite"); err != nil {
|
||||
t.Fatalf("register sharding: %v", err)
|
||||
}
|
||||
if err := autoMigrateAll(db); err != nil {
|
||||
t.Fatalf("auto migrate db: %v", err)
|
||||
}
|
||||
|
||||
for _, table := range []string{
|
||||
"node_metric_snapshots_00",
|
||||
"node_metric_snapshots_09",
|
||||
"node_request_reports_00",
|
||||
"node_request_reports_09",
|
||||
"node_access_logs_00",
|
||||
"node_access_logs_09",
|
||||
} {
|
||||
if !db.Migrator().HasTable(table) {
|
||||
t.Fatalf("expected sharded table %s to exist", table)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpgradeDatabaseSchemaV15ToV16AppliesCompressedReleaseSchema(t *testing.T) {
|
||||
db := openBareTestSQLiteDB(t, "v16.db")
|
||||
if err := registerSharding(db, "sqlite"); err != nil {
|
||||
t.Fatalf("register sharding: %v", err)
|
||||
}
|
||||
if err := autoMigrateLegacySchemaMetadata(db); err != nil {
|
||||
t.Fatalf("auto migrate legacy schema metadata: %v", err)
|
||||
}
|
||||
if err := applyCurrentSchema(db, "sqlite"); err != nil {
|
||||
t.Fatalf("apply current schema: %v", err)
|
||||
}
|
||||
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
|
||||
t.Fatalf("failed to add legacy pow_enabled: %v", err)
|
||||
}
|
||||
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
|
||||
t.Fatalf("failed to add legacy pow_config: %v", err)
|
||||
}
|
||||
if err := ensureDefaultWAFRuleGroup(db); err != nil {
|
||||
t.Fatalf("ensure default waf rule group: %v", err)
|
||||
}
|
||||
if err := saveDatabaseSchemaVersion(db, 15); err != nil {
|
||||
t.Fatalf("save schema version: %v", err)
|
||||
}
|
||||
if err := upgradeDatabaseSchema(db, "sqlite", 15); err != nil {
|
||||
t.Fatalf("upgrade schema: %v", err)
|
||||
}
|
||||
if !db.Migrator().HasTable(&WAFIPGroup{}) {
|
||||
t.Fatal("expected waf_ip_groups table")
|
||||
}
|
||||
if !db.Migrator().HasColumn(&WAFRuleGroup{}, "ip_whitelist_groups") {
|
||||
t.Fatal("expected waf_rule_groups.ip_whitelist_groups column")
|
||||
}
|
||||
if !db.Migrator().HasColumn(&Node{}, "access_token") {
|
||||
t.Fatal("expected nodes.access_token column")
|
||||
}
|
||||
if !db.Migrator().HasColumn(&Node{}, "version") {
|
||||
t.Fatal("expected nodes.version column")
|
||||
}
|
||||
if !db.Migrator().HasColumn(&Node{}, "ext_version") {
|
||||
t.Fatal("expected nodes.ext_version column")
|
||||
}
|
||||
if !db.Migrator().HasColumn(&ProxyRoute{}, "tunnel_node_id") {
|
||||
t.Fatal("expected proxy_routes.tunnel_node_id column")
|
||||
}
|
||||
if db.Migrator().HasTable("tunnels") {
|
||||
t.Fatal("expected pre-release tunnels table to be absent")
|
||||
}
|
||||
version, ok, err := loadDatabaseSchemaVersion(db)
|
||||
if err != nil {
|
||||
t.Fatalf("load schema version: %v", err)
|
||||
}
|
||||
if !ok || version != currentDatabaseSchemaVersion {
|
||||
t.Fatalf("unexpected schema version: got %d ok=%v want %d", version, ok, currentDatabaseSchemaVersion)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrateObservabilityLegacyColumnsBackfillsHealthEventMetadata(t *testing.T) {
|
||||
db := openTestSQLiteDB(t, "legacy-health-events.db")
|
||||
|
||||
if err := db.Exec("ALTER TABLE node_health_events ADD COLUMN raw_json TEXT").Error; err != nil {
|
||||
t.Fatalf("add raw_json column: %v", err)
|
||||
}
|
||||
rawJSON, err := json.Marshal(map[string]any{
|
||||
"event_type": "sync_error",
|
||||
"metadata": map[string]string{
|
||||
"reason": "checksum_mismatch",
|
||||
"scope": "routes",
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal raw json: %v", err)
|
||||
}
|
||||
event := &NodeHealthEvent{
|
||||
NodeID: "node-legacy",
|
||||
EventType: "sync_error",
|
||||
Severity: "warning",
|
||||
Status: "active",
|
||||
Message: "checksum mismatch",
|
||||
FirstTriggeredAt: time.Now().Add(-time.Minute),
|
||||
LastTriggeredAt: time.Now(),
|
||||
ReportedAt: time.Now(),
|
||||
}
|
||||
if err := db.Create(event).Error; err != nil {
|
||||
t.Fatalf("create health event: %v", err)
|
||||
}
|
||||
if err := db.Exec("UPDATE node_health_events SET raw_json = ? WHERE id = ?", string(rawJSON), event.ID).Error; err != nil {
|
||||
t.Fatalf("seed legacy raw_json: %v", err)
|
||||
}
|
||||
|
||||
if err := migrateObservabilityLegacyColumns(db); err != nil {
|
||||
t.Fatalf("migrateObservabilityLegacyColumns: %v", err)
|
||||
}
|
||||
|
||||
var got NodeHealthEvent
|
||||
if err := db.First(&got, event.ID).Error; err != nil {
|
||||
t.Fatalf("query health event: %v", err)
|
||||
}
|
||||
if got.MetadataJSON == "" {
|
||||
t.Fatal("expected metadata_json to be backfilled")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureDatabaseSchemaUpToDateInitializesFreshDatabase(t *testing.T) {
|
||||
db := openBareTestSQLiteDB(t, "fresh-schema.db")
|
||||
if err := registerSharding(db, "sqlite"); err != nil {
|
||||
t.Fatalf("register sharding: %v", err)
|
||||
}
|
||||
|
||||
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
|
||||
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
|
||||
}
|
||||
|
||||
version, exists, err := loadDatabaseSchemaVersion(db)
|
||||
if err != nil {
|
||||
t.Fatalf("loadDatabaseSchemaVersion: %v", err)
|
||||
}
|
||||
if !exists {
|
||||
t.Fatal("expected database schema version to be recorded")
|
||||
}
|
||||
if version != expectedCurrentDatabaseVersion() {
|
||||
t.Fatalf("unexpected schema version: got %d want %d", version, expectedCurrentDatabaseVersion())
|
||||
}
|
||||
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
|
||||
t.Fatal("expected fresh database to avoid legacy database_schema_versions table")
|
||||
}
|
||||
if !db.Migrator().HasTable("goose_db_version") {
|
||||
t.Fatal("expected fresh database to initialize goose_db_version")
|
||||
}
|
||||
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
|
||||
t.Fatal("expected fresh database to apply goose migration nodes.capabilities_json")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureDatabaseSchemaUpToDateUpgradesLegacyDatabase(t *testing.T) {
|
||||
db := openBareTestSQLiteDB(t, "legacy-schema.db")
|
||||
if err := registerSharding(db, "sqlite"); err != nil {
|
||||
t.Fatalf("register sharding: %v", err)
|
||||
}
|
||||
if err := autoMigrateAll(db); err != nil {
|
||||
t.Fatalf("auto migrate db: %v", err)
|
||||
}
|
||||
// Add legacy PoW columns manually to proxy_routes table to simulate legacy schema v9-v17 state
|
||||
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
|
||||
t.Fatalf("failed to add legacy pow_enabled: %v", err)
|
||||
}
|
||||
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
|
||||
t.Fatalf("failed to add legacy pow_config: %v", err)
|
||||
}
|
||||
if err := db.Create(&User{
|
||||
Username: "legacy",
|
||||
Password: "secret",
|
||||
DisplayName: "Legacy User",
|
||||
Role: 1,
|
||||
Status: 1,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed legacy user: %v", err)
|
||||
}
|
||||
|
||||
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
|
||||
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
|
||||
}
|
||||
|
||||
version, exists, err := loadDatabaseSchemaVersion(db)
|
||||
if err != nil {
|
||||
t.Fatalf("loadDatabaseSchemaVersion: %v", err)
|
||||
}
|
||||
if !exists {
|
||||
t.Fatal("expected legacy database to gain a schema version record")
|
||||
}
|
||||
if version != expectedCurrentDatabaseVersion() {
|
||||
t.Fatalf("unexpected schema version: got %d want %d", version, expectedCurrentDatabaseVersion())
|
||||
}
|
||||
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
|
||||
t.Fatal("expected legacy database_schema_versions table to be removed after bridging to goose")
|
||||
}
|
||||
if !db.Migrator().HasTable("goose_db_version") {
|
||||
t.Fatal("expected legacy upgrade to initialize goose_db_version")
|
||||
}
|
||||
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
|
||||
t.Fatal("expected legacy upgrade to apply goose migration nodes.capabilities_json")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrateOriginsSchemaBackfillsOrigins(t *testing.T) {
|
||||
db := openBareTestSQLiteDB(t, "legacy-origins.db")
|
||||
if err := registerSharding(db, "sqlite"); err != nil {
|
||||
t.Fatalf("register sharding: %v", err)
|
||||
}
|
||||
if err := applyCurrentSchema(db, "sqlite"); err != nil {
|
||||
t.Fatalf("applyCurrentSchema: %v", err)
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
route := &ProxyRoute{
|
||||
Domain: "app.example.com",
|
||||
OriginURL: "https://origin-a.internal:8443/api",
|
||||
Upstreams: `["https://origin-a.internal:8443/api"]`,
|
||||
Enabled: true,
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
}
|
||||
if err := db.Create(route).Error; err != nil {
|
||||
t.Fatalf("seed proxy route: %v", err)
|
||||
}
|
||||
if err := db.Exec(`DELETE FROM origins`).Error; err != nil {
|
||||
t.Fatalf("clear origins: %v", err)
|
||||
}
|
||||
if err := db.Model(&ProxyRoute{}).Where("id = ?", route.ID).Update("origin_id", nil).Error; err != nil {
|
||||
t.Fatalf("clear route origin_id: %v", err)
|
||||
}
|
||||
|
||||
if err := backfillOriginsFromProxyRoutes(db); err != nil {
|
||||
t.Fatalf("backfillOriginsFromProxyRoutes: %v", err)
|
||||
}
|
||||
|
||||
if !db.Migrator().HasTable(&Origin{}) {
|
||||
t.Fatal("expected origins table to exist")
|
||||
}
|
||||
if !db.Migrator().HasColumn(&ProxyRoute{}, "origin_id") {
|
||||
t.Fatal("expected proxy_routes.origin_id column to exist")
|
||||
}
|
||||
|
||||
reloadedRoute := &ProxyRoute{}
|
||||
if err := db.First(reloadedRoute, route.ID).Error; err != nil {
|
||||
t.Fatalf("query proxy route: %v", err)
|
||||
}
|
||||
if reloadedRoute.OriginID == nil || *reloadedRoute.OriginID == 0 {
|
||||
t.Fatal("expected migrated route to be linked to a backfilled origin")
|
||||
}
|
||||
|
||||
origin := &Origin{}
|
||||
if err := db.First(origin, *reloadedRoute.OriginID).Error; err != nil {
|
||||
t.Fatalf("query origin: %v", err)
|
||||
}
|
||||
if origin.Address != "origin-a.internal" {
|
||||
t.Fatalf("unexpected backfilled origin address: %s", origin.Address)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureDatabaseSchemaUpToDateAddsProxyRouteDomainCertificateFields(t *testing.T) {
|
||||
db := openBareTestSQLiteDB(t, "legacy-proxy-route-domain-cert-ids.db")
|
||||
if err := registerSharding(db, "sqlite"); err != nil {
|
||||
t.Fatalf("register sharding: %v", err)
|
||||
}
|
||||
if err := autoMigrateLegacySchemaMetadata(db); err != nil {
|
||||
t.Fatalf("auto migrate legacy schema metadata: %v", err)
|
||||
}
|
||||
|
||||
for _, item := range registeredModels() {
|
||||
if _, ok := item.(*ProxyRoute); ok {
|
||||
continue
|
||||
}
|
||||
if err := db.AutoMigrate(item); err != nil {
|
||||
t.Fatalf("auto migrate supporting table: %v", err)
|
||||
}
|
||||
}
|
||||
if err := db.AutoMigrate(&legacyProxyRouteV7{}); err != nil {
|
||||
t.Fatalf("auto migrate legacy proxy_routes v7: %v", err)
|
||||
}
|
||||
// Add legacy PoW columns manually to proxy_routes table to simulate legacy schema v9-v17 state
|
||||
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
|
||||
t.Fatalf("failed to add legacy pow_enabled: %v", err)
|
||||
}
|
||||
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
|
||||
t.Fatalf("failed to add legacy pow_config: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now().UTC()
|
||||
certID := uint(9)
|
||||
if err := db.Create(&legacyProxyRouteV7{
|
||||
SiteName: "secure-site",
|
||||
Domain: "secure.example.com",
|
||||
Domains: `["secure.example.com","www.secure.example.com"]`,
|
||||
OriginURL: "https://origin-secure.internal:8443",
|
||||
Upstreams: `["https://origin-secure.internal:8443"]`,
|
||||
Enabled: true,
|
||||
EnableHTTPS: true,
|
||||
CertID: &certID,
|
||||
CertIDs: `[9]`,
|
||||
RedirectHTTP: true,
|
||||
LimitConnPerServer: 120,
|
||||
LimitConnPerIP: 12,
|
||||
LimitRate: "512k",
|
||||
CacheEnabled: false,
|
||||
CachePolicy: "",
|
||||
CacheRules: `[]`,
|
||||
CustomHeaders: `[]`,
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed legacy proxy route v7: %v", err)
|
||||
}
|
||||
if err := saveDatabaseSchemaVersion(db, 7); err != nil {
|
||||
t.Fatalf("save schema version: %v", err)
|
||||
}
|
||||
|
||||
previousDB := DB
|
||||
DB = db
|
||||
t.Cleanup(func() {
|
||||
DB = previousDB
|
||||
})
|
||||
|
||||
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
|
||||
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
|
||||
}
|
||||
|
||||
var route ProxyRoute
|
||||
if err := db.First(&route).Error; err != nil {
|
||||
t.Fatalf("query migrated proxy route: %v", err)
|
||||
}
|
||||
|
||||
var domainCertIDs []uint
|
||||
if err := json.Unmarshal([]byte(route.DomainCertIDs), &domainCertIDs); err != nil {
|
||||
t.Fatalf("decode migrated domain_cert_ids: %v", err)
|
||||
}
|
||||
if len(domainCertIDs) != 2 || domainCertIDs[0] != certID || domainCertIDs[1] != certID {
|
||||
t.Fatalf("unexpected migrated domain_cert_ids: %#v", domainCertIDs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunDatabaseSchemaMigrationDoesNotAdvanceVersionWhenValidationFails(t *testing.T) {
|
||||
db := openBareTestSQLiteDB(t, "failed-validation.db")
|
||||
|
||||
err := runDatabaseSchemaMigration(db, "sqlite", databaseSchemaMigration{
|
||||
fromVersion: legacyDatabaseSchemaVersion,
|
||||
toVersion: 11,
|
||||
migrate: func(tx *gorm.DB, backend string) error {
|
||||
return autoMigrateLegacySchemaMetadata(tx)
|
||||
},
|
||||
validate: func(tx *gorm.DB, backend string) error {
|
||||
return gorm.ErrInvalidDB
|
||||
},
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected migration validation to fail")
|
||||
}
|
||||
|
||||
_, exists, loadErr := loadDatabaseSchemaVersion(db)
|
||||
if loadErr != nil {
|
||||
t.Fatalf("loadDatabaseSchemaVersion: %v", loadErr)
|
||||
}
|
||||
if exists {
|
||||
t.Fatal("expected schema version to remain unset after failed validation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureDatabaseSchemaUpToDateAddsNodeIPManualOverride(t *testing.T) {
|
||||
db := openBareTestSQLiteDB(t, "node-ip-manual-override-migration.db")
|
||||
if err := registerSharding(db, "sqlite"); err != nil {
|
||||
t.Fatalf("register sharding: %v", err)
|
||||
}
|
||||
if err := applyCurrentSchema(db, "sqlite"); err != nil {
|
||||
t.Fatalf("apply current schema: %v", err)
|
||||
}
|
||||
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
|
||||
t.Fatalf("failed to add legacy pow_enabled: %v", err)
|
||||
}
|
||||
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
|
||||
t.Fatalf("failed to add legacy pow_config: %v", err)
|
||||
}
|
||||
if err := ensureDefaultWAFRuleGroup(db); err != nil {
|
||||
t.Fatalf("ensure default waf rule group: %v", err)
|
||||
}
|
||||
if err := db.Migrator().DropColumn(&Node{}, "ip_manual_override"); err != nil {
|
||||
t.Fatalf("drop ip_manual_override column: %v", err)
|
||||
}
|
||||
if db.Migrator().HasColumn(&Node{}, "ip_manual_override") {
|
||||
t.Fatal("expected test database to simulate schema v14 without ip_manual_override")
|
||||
}
|
||||
if err := saveDatabaseSchemaVersion(db, 14); err != nil {
|
||||
t.Fatalf("save schema version: %v", err)
|
||||
}
|
||||
|
||||
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
|
||||
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
|
||||
}
|
||||
|
||||
if !db.Migrator().HasColumn(&Node{}, "ip_manual_override") {
|
||||
t.Fatal("expected migration to add nodes.ip_manual_override")
|
||||
}
|
||||
version, exists, err := loadDatabaseSchemaVersion(db)
|
||||
if err != nil {
|
||||
t.Fatalf("loadDatabaseSchemaVersion: %v", err)
|
||||
}
|
||||
if !exists {
|
||||
t.Fatal("expected schema version record to exist")
|
||||
}
|
||||
if version != expectedCurrentDatabaseVersion() {
|
||||
t.Fatalf("unexpected schema version: got %d want %d", version, expectedCurrentDatabaseVersion())
|
||||
}
|
||||
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
|
||||
t.Fatal("expected migration chain to include nodes.capabilities_json")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureDatabaseSchemaUpToDateV16BackfillsNodeColumnsWhenNewColumnsAlreadyExist(t *testing.T) {
|
||||
db := openBareTestSQLiteDB(t, "node-v16-existing-target-columns.db")
|
||||
if err := registerSharding(db, "sqlite"); err != nil {
|
||||
t.Fatalf("register sharding: %v", err)
|
||||
}
|
||||
if err := applyCurrentSchema(db, "sqlite"); err != nil {
|
||||
t.Fatalf("apply current schema: %v", err)
|
||||
}
|
||||
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
|
||||
t.Fatalf("failed to add legacy pow_enabled: %v", err)
|
||||
}
|
||||
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
|
||||
t.Fatalf("failed to add legacy pow_config: %v", err)
|
||||
}
|
||||
if err := ensureDefaultWAFRuleGroup(db); err != nil {
|
||||
t.Fatalf("ensure default waf rule group: %v", err)
|
||||
}
|
||||
for _, stmt := range []string{
|
||||
`ALTER TABLE nodes ADD COLUMN agent_token text`,
|
||||
`ALTER TABLE nodes ADD COLUMN agent_version text`,
|
||||
`ALTER TABLE nodes ADD COLUMN nginx_version text`,
|
||||
`ALTER TABLE nodes ADD COLUMN relay_version text`,
|
||||
`ALTER TABLE nodes ADD COLUMN relay_frp_version text`,
|
||||
`ALTER TABLE nodes ADD COLUMN relay_frps_connections integer`,
|
||||
`ALTER TABLE nodes ADD COLUMN relay_frps_proxy_count integer`,
|
||||
} {
|
||||
if err := db.Exec(stmt).Error; err != nil {
|
||||
t.Fatalf("prepare legacy node column with %q: %v", stmt, err)
|
||||
}
|
||||
}
|
||||
now := time.Now()
|
||||
if err := db.Exec(`
|
||||
INSERT INTO nodes (
|
||||
node_id, name, ip, access_token, version, ext_version,
|
||||
agent_token, agent_version, nginx_version,
|
||||
status, last_seen_at, created_at, updated_at
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "node-v16", "Node v16", "127.0.0.1", "", "", "", "legacy-token", "v2.0.0", "openresty/1.25.3", "offline", now, now, now).Error; err != nil {
|
||||
t.Fatalf("seed node with legacy columns: %v", err)
|
||||
}
|
||||
if err := saveDatabaseSchemaVersion(db, 15); err != nil {
|
||||
t.Fatalf("save schema version: %v", err)
|
||||
}
|
||||
|
||||
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
|
||||
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
|
||||
}
|
||||
|
||||
var node Node
|
||||
if err := db.Where("node_id = ?", "node-v16").First(&node).Error; err != nil {
|
||||
t.Fatalf("query migrated node: %v", err)
|
||||
}
|
||||
if node.AccessToken != "legacy-token" {
|
||||
t.Fatalf("unexpected access_token: got %q", node.AccessToken)
|
||||
}
|
||||
if node.Version != "v2.0.0" {
|
||||
t.Fatalf("unexpected version: got %q", node.Version)
|
||||
}
|
||||
if node.ExtVersion != "openresty/1.25.3" {
|
||||
t.Fatalf("unexpected ext_version: got %q", node.ExtVersion)
|
||||
}
|
||||
for _, column := range []string{
|
||||
"agent_token",
|
||||
"agent_version",
|
||||
"nginx_version",
|
||||
"relay_version",
|
||||
"relay_frp_version",
|
||||
"relay_frps_connections",
|
||||
"relay_frps_proxy_count",
|
||||
} {
|
||||
exists, err := databaseColumnExists(db, "nodes", column)
|
||||
if err != nil {
|
||||
t.Fatalf("inspect legacy nodes.%s: %v", column, err)
|
||||
}
|
||||
if exists {
|
||||
t.Fatalf("expected migration to drop legacy nodes.%s column", column)
|
||||
}
|
||||
}
|
||||
version, exists, err := loadDatabaseSchemaVersion(db)
|
||||
if err != nil {
|
||||
t.Fatalf("loadDatabaseSchemaVersion: %v", err)
|
||||
}
|
||||
if !exists {
|
||||
t.Fatal("expected schema version record to exist")
|
||||
}
|
||||
if version != expectedCurrentDatabaseVersion() {
|
||||
t.Fatalf("unexpected schema version: got %d want %d", version, expectedCurrentDatabaseVersion())
|
||||
}
|
||||
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
|
||||
t.Fatal("expected v16 upgrade path to apply goose migration nodes.capabilities_json")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureDatabaseSchemaUpToDateV16DropsLegacyNodeColumnsWhenAlreadyCurrent(t *testing.T) {
|
||||
db := openBareTestSQLiteDB(t, "node-v16-current-legacy-columns.db")
|
||||
if err := registerSharding(db, "sqlite"); err != nil {
|
||||
t.Fatalf("register sharding: %v", err)
|
||||
}
|
||||
if err := applyCurrentSchema(db, "sqlite"); err != nil {
|
||||
t.Fatalf("apply current schema: %v", err)
|
||||
}
|
||||
for _, stmt := range []string{
|
||||
`ALTER TABLE nodes ADD COLUMN agent_token text`,
|
||||
`ALTER TABLE nodes ADD COLUMN agent_version text`,
|
||||
`ALTER TABLE nodes ADD COLUMN nginx_version text`,
|
||||
} {
|
||||
if err := db.Exec(stmt).Error; err != nil {
|
||||
t.Fatalf("prepare legacy node column with %q: %v", stmt, err)
|
||||
}
|
||||
}
|
||||
if err := saveDatabaseSchemaVersion(db, currentDatabaseSchemaVersion); err != nil {
|
||||
t.Fatalf("save schema version: %v", err)
|
||||
}
|
||||
|
||||
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
|
||||
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
|
||||
}
|
||||
|
||||
for _, column := range []string{"agent_token", "agent_version", "nginx_version"} {
|
||||
exists, err := databaseColumnExists(db, "nodes", column)
|
||||
if err != nil {
|
||||
t.Fatalf("inspect legacy nodes.%s: %v", column, err)
|
||||
}
|
||||
if exists {
|
||||
t.Fatalf("expected current-schema cleanup to drop legacy nodes.%s column", column)
|
||||
}
|
||||
}
|
||||
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
|
||||
t.Fatal("expected current-schema legacy version table to be removed after goose bridge")
|
||||
}
|
||||
if !db.Migrator().HasTable("goose_db_version") {
|
||||
t.Fatal("expected current-schema goose_db_version table to exist")
|
||||
}
|
||||
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
|
||||
t.Fatal("expected current-schema repair to preserve goose column nodes.capabilities_json")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureDatabaseSchemaUpToDateKeepsGooseOnlyDatabaseOnReentry(t *testing.T) {
|
||||
db := openBareTestSQLiteDB(t, "goose-only-reentry.db")
|
||||
if err := registerSharding(db, "sqlite"); err != nil {
|
||||
t.Fatalf("register sharding: %v", err)
|
||||
}
|
||||
|
||||
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
|
||||
t.Fatalf("first ensureDatabaseSchemaUpToDate: %v", err)
|
||||
}
|
||||
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
|
||||
t.Fatal("expected first initialization to avoid legacy table")
|
||||
}
|
||||
|
||||
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
|
||||
t.Fatalf("second ensureDatabaseSchemaUpToDate: %v", err)
|
||||
}
|
||||
|
||||
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
|
||||
t.Fatal("expected goose-only database to remain free of legacy version table")
|
||||
}
|
||||
version, exists, err := loadGooseDatabaseVersion(db)
|
||||
if err != nil {
|
||||
t.Fatalf("loadGooseDatabaseVersion: %v", err)
|
||||
}
|
||||
if !exists {
|
||||
t.Fatal("expected goose-only database to keep goose version record")
|
||||
}
|
||||
if version != expectedCurrentDatabaseVersion() {
|
||||
t.Fatalf("unexpected goose version: got %d want %d", version, expectedCurrentDatabaseVersion())
|
||||
}
|
||||
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
|
||||
t.Fatal("expected goose-only database to keep nodes.capabilities_json")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAllRegisteredMigrationsHaveValidationDefined(t *testing.T) {
|
||||
ctx := databaseSchemaMigrationContext{}
|
||||
for _, migration := range databaseSchemaMigrations() {
|
||||
err := ctx.ValidateDatabaseSchemaVersion(nil, "sqlite", migration.toVersion)
|
||||
if err != nil && strings.Contains(err.Error(), "is not defined") {
|
||||
t.Fatalf("Validation is not defined in migrations.go for registered migration version v%d: %v", migration.toVersion, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAllGORMModelsAreRegistered(t *testing.T) {
|
||||
// 1. Gather all registered model names
|
||||
registeredNames := make(map[string]bool)
|
||||
for _, item := range registeredModels() {
|
||||
name := reflect.TypeOf(item).Elem().Name()
|
||||
registeredNames[name] = true
|
||||
}
|
||||
for _, item := range schemaMetadataModels() {
|
||||
name := reflect.TypeOf(item).Elem().Name()
|
||||
registeredNames[name] = true
|
||||
}
|
||||
|
||||
// 2. Parse all .go files in model/ package
|
||||
fset := token.NewFileSet()
|
||||
pkgs, err := parser.ParseDir(fset, ".", func(info os.FileInfo) bool {
|
||||
// Only parse .go files, exclude _test.go files and subdirectories
|
||||
return !info.IsDir() && strings.HasSuffix(info.Name(), ".go") && !strings.HasSuffix(info.Name(), "_test.go")
|
||||
}, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to parse directory: %v", err)
|
||||
}
|
||||
|
||||
for _, pkg := range pkgs {
|
||||
for _, file := range pkg.Files {
|
||||
for _, decl := range file.Decls {
|
||||
genDecl, ok := decl.(*ast.GenDecl)
|
||||
if !ok || genDecl.Tok != token.TYPE {
|
||||
continue
|
||||
}
|
||||
for _, spec := range genDecl.Specs {
|
||||
typeSpec, ok := spec.(*ast.TypeSpec)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
structType, ok := typeSpec.Type.(*ast.StructType)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
|
||||
// Verify if this struct has any field with a `gorm:"..."` tag
|
||||
isGORMModel := false
|
||||
for _, field := range structType.Fields.List {
|
||||
if field.Tag != nil && strings.Contains(field.Tag.Value, "gorm:") {
|
||||
isGORMModel = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if isGORMModel {
|
||||
structName := typeSpec.Name.Name
|
||||
if !registeredNames[structName] {
|
||||
t.Errorf("Model struct %q is defined with GORM tags but is NOT registered in registeredModels() or schemaMetadataModels() in model/main.go!", structName)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureDatabaseSchemaUpToDateDropsPagesDeploymentUnusedFields(t *testing.T) {
|
||||
db := openBareTestSQLiteDB(t, "drop-pages-deployment-unused-fields.db")
|
||||
if err := registerSharding(db, "sqlite"); err != nil {
|
||||
t.Fatalf("register sharding: %v", err)
|
||||
}
|
||||
|
||||
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
|
||||
t.Fatalf("first ensureDatabaseSchemaUpToDate: %v", err)
|
||||
}
|
||||
|
||||
// Verify columns do not exist
|
||||
if db.Migrator().HasColumn("pages_deployments", "root_dir") {
|
||||
t.Fatal("expected root_dir column to be absent initially")
|
||||
}
|
||||
if db.Migrator().HasColumn("pages_deployments", "entry_file") {
|
||||
t.Fatal("expected entry_file column to be absent initially")
|
||||
}
|
||||
|
||||
// Manually add columns to simulate old state
|
||||
if err := db.Exec("ALTER TABLE pages_deployments ADD COLUMN root_dir TEXT").Error; err != nil {
|
||||
t.Fatalf("failed to add root_dir column: %v", err)
|
||||
}
|
||||
if err := db.Exec("ALTER TABLE pages_deployments ADD COLUMN entry_file TEXT").Error; err != nil {
|
||||
t.Fatalf("failed to add entry_file column: %v", err)
|
||||
}
|
||||
|
||||
// Verify columns were added
|
||||
if !db.Migrator().HasColumn("pages_deployments", "root_dir") {
|
||||
t.Fatal("expected root_dir column to be present after manual add")
|
||||
}
|
||||
if !db.Migrator().HasColumn("pages_deployments", "entry_file") {
|
||||
t.Fatal("expected entry_file column to be present after manual add")
|
||||
}
|
||||
|
||||
// Remove the migration record from goose_db_version table
|
||||
const versionToRerun = 202606040004
|
||||
if err := db.Exec("DELETE FROM goose_db_version WHERE version_id = ?", versionToRerun).Error; err != nil {
|
||||
t.Fatalf("failed to delete migration record: %v", err)
|
||||
}
|
||||
|
||||
// Run migration again
|
||||
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
|
||||
t.Fatalf("second ensureDatabaseSchemaUpToDate: %v", err)
|
||||
}
|
||||
|
||||
// Verify columns were dropped successfully
|
||||
if db.Migrator().HasColumn("pages_deployments", "root_dir") {
|
||||
t.Fatal("expected root_dir column to be dropped after migration rerun")
|
||||
}
|
||||
if db.Migrator().HasColumn("pages_deployments", "entry_file") {
|
||||
t.Fatal("expected entry_file column to be dropped after migration rerun")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type ManagedDomain struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
Domain string `json:"domain" gorm:"uniqueIndex;size:255;not null"`
|
||||
CertID *uint `json:"cert_id"`
|
||||
Enabled bool `json:"enabled" gorm:"not null;default:true"`
|
||||
Remark string `json:"remark" gorm:"size:255"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
func ListManagedDomains() (domains []*ManagedDomain, err error) {
|
||||
err = DB.Order("id desc").Find(&domains).Error
|
||||
return domains, err
|
||||
}
|
||||
|
||||
func ListEnabledManagedDomainsWithCertificate() (domains []*ManagedDomain, err error) {
|
||||
err = DB.Where("enabled = ? AND cert_id IS NOT NULL", true).Order("id desc").Find(&domains).Error
|
||||
return domains, err
|
||||
}
|
||||
|
||||
func GetManagedDomainByID(id uint) (*ManagedDomain, error) {
|
||||
domain := &ManagedDomain{}
|
||||
err := DB.First(domain, id).Error
|
||||
return domain, err
|
||||
}
|
||||
|
||||
func (domain *ManagedDomain) Insert() error {
|
||||
return DB.Create(domain).Error
|
||||
}
|
||||
|
||||
func (domain *ManagedDomain) Update() error {
|
||||
return DB.Save(domain).Error
|
||||
}
|
||||
|
||||
func (domain *ManagedDomain) Delete() error {
|
||||
return DB.Delete(domain).Error
|
||||
}
|
||||
@@ -0,0 +1,4 @@
|
||||
package migrate
|
||||
|
||||
// Versions 1 through 7 are treated as the historical baseline. There are no
|
||||
// supported deployments below v8, so new upgrades start from this base version.
|
||||
@@ -0,0 +1,54 @@
|
||||
package migrate
|
||||
|
||||
import (
|
||||
"sort"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const BaseDatabaseSchemaVersion = 7
|
||||
|
||||
type Context interface {
|
||||
ApplyCurrentSchema(db *gorm.DB, backend string) error
|
||||
ApplyCurrentSchemaExcept(db *gorm.DB, backend string, excludedTables ...string) error
|
||||
BackfillOriginsFromProxyRoutes(db *gorm.DB) error
|
||||
BackfillProxyRouteSiteFields(db *gorm.DB) error
|
||||
EnsureProxyRouteSiteNameUniqueIndex(db *gorm.DB) error
|
||||
BackfillProxyRouteCertificateFields(db *gorm.DB) error
|
||||
BackfillProxyRouteDomainCertificateFields(db *gorm.DB) error
|
||||
EnsureDefaultGitHubAuthSource(db *gorm.DB) error
|
||||
EnsureDefaultWAFRuleGroup(db *gorm.DB) error
|
||||
DropLegacyNodeColumns(db *gorm.DB, backend string) error
|
||||
ValidateDatabaseSchemaVersion(db *gorm.DB, backend string, version int) error
|
||||
}
|
||||
|
||||
type Migration struct {
|
||||
FromVersion int
|
||||
ToVersion int
|
||||
Migrate func(ctx Context, db *gorm.DB, backend string) error
|
||||
Validate func(ctx Context, db *gorm.DB, backend string) error
|
||||
}
|
||||
|
||||
var registeredMigrations []Migration
|
||||
|
||||
func Register(migration Migration) {
|
||||
registeredMigrations = append(registeredMigrations, migration)
|
||||
}
|
||||
|
||||
func Migrations() []Migration {
|
||||
migrations := append([]Migration{}, registeredMigrations...)
|
||||
sort.Slice(migrations, func(i int, j int) bool {
|
||||
return migrations[i].FromVersion < migrations[j].FromVersion
|
||||
})
|
||||
return migrations
|
||||
}
|
||||
|
||||
func CurrentVersion() int {
|
||||
version := BaseDatabaseSchemaVersion
|
||||
for _, migration := range registeredMigrations {
|
||||
if migration.ToVersion > version {
|
||||
version = migration.ToVersion
|
||||
}
|
||||
}
|
||||
return version
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
package migrate
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestMigrationsAreContinuousFromBaseVersion(t *testing.T) {
|
||||
migrations := Migrations()
|
||||
if len(migrations) == 0 {
|
||||
t.Fatal("expected at least one registered migration")
|
||||
}
|
||||
expectedFrom := BaseDatabaseSchemaVersion
|
||||
for _, migration := range migrations {
|
||||
if migration.FromVersion != expectedFrom {
|
||||
t.Fatalf("expected migration from v%d, got v%d -> v%d", expectedFrom, migration.FromVersion, migration.ToVersion)
|
||||
}
|
||||
if migration.ToVersion != migration.FromVersion+1 {
|
||||
t.Fatalf("expected one-step migration, got v%d -> v%d", migration.FromVersion, migration.ToVersion)
|
||||
}
|
||||
expectedFrom = migration.ToVersion
|
||||
}
|
||||
if CurrentVersion() != expectedFrom {
|
||||
t.Fatalf("unexpected current version: got %d want %d", CurrentVersion(), expectedFrom)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
// v10 升级内容:新增可配置认证源与第三方账号绑定,并迁移旧 GitHub 登录配置。
|
||||
// 背景说明:登录体系从固定 GitHub OAuth 字段演进为通用认证源模型,需要创建 auth_sources、external_accounts,并把旧用户 GitHub 绑定迁移到新表。
|
||||
package migrate
|
||||
|
||||
import "gorm.io/gorm"
|
||||
|
||||
func init() {
|
||||
Register(V10())
|
||||
}
|
||||
|
||||
func V10() Migration {
|
||||
return Migration{
|
||||
FromVersion: 9,
|
||||
ToVersion: 10,
|
||||
Migrate: migrateV10,
|
||||
Validate: validateV10,
|
||||
}
|
||||
}
|
||||
|
||||
func migrateV10(ctx Context, db *gorm.DB, backend string) error {
|
||||
return ctx.EnsureDefaultGitHubAuthSource(db)
|
||||
}
|
||||
|
||||
func validateV10(ctx Context, db *gorm.DB, backend string) error {
|
||||
return ctx.ValidateDatabaseSchemaVersion(db, backend, 10)
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
// v11 升级内容:新增 ACME 账户、DNS 账户,并扩展证书 provider 字段。
|
||||
// 背景说明:证书申请能力从单一手工导入扩展到自动签发,需要持久化 ACME/DNS 凭据,并标记证书来源。
|
||||
package migrate
|
||||
|
||||
import "gorm.io/gorm"
|
||||
|
||||
func init() {
|
||||
Register(V11())
|
||||
}
|
||||
|
||||
func V11() Migration {
|
||||
return Migration{
|
||||
FromVersion: 10,
|
||||
ToVersion: 11,
|
||||
Migrate: migrateV11,
|
||||
Validate: validateV11,
|
||||
}
|
||||
}
|
||||
|
||||
func migrateV11(ctx Context, db *gorm.DB, backend string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateV11(ctx Context, db *gorm.DB, backend string) error {
|
||||
return ctx.ValidateDatabaseSchemaVersion(db, backend, 11)
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
// v12 升级内容:为 proxy_routes 增加 Basic Auth 相关字段。
|
||||
// 背景说明:站点级访问控制需要支持基础认证,因此在代理路由配置中持久化 Basic Auth 开关与凭据配置。
|
||||
package migrate
|
||||
|
||||
import "gorm.io/gorm"
|
||||
|
||||
func init() {
|
||||
Register(V12())
|
||||
}
|
||||
|
||||
func V12() Migration {
|
||||
return Migration{
|
||||
FromVersion: 11,
|
||||
ToVersion: 12,
|
||||
Migrate: migrateV12,
|
||||
Validate: validateV12,
|
||||
}
|
||||
}
|
||||
|
||||
func migrateV12(ctx Context, db *gorm.DB, backend string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateV12(ctx Context, db *gorm.DB, backend string) error {
|
||||
return ctx.ValidateDatabaseSchemaVersion(db, backend, 12)
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
// v13 升级内容:新增 WAF 规则组与站点绑定表,并创建默认全局规则组。
|
||||
// 背景说明:WAF 配置从零散站点字段演进为可复用规则组,需要全局规则组作为默认入口,并支持站点与规则组绑定。
|
||||
package migrate
|
||||
|
||||
import "gorm.io/gorm"
|
||||
|
||||
func init() {
|
||||
Register(V13())
|
||||
}
|
||||
|
||||
func V13() Migration {
|
||||
return Migration{
|
||||
FromVersion: 12,
|
||||
ToVersion: 13,
|
||||
Migrate: migrateV13,
|
||||
Validate: validateV13,
|
||||
}
|
||||
}
|
||||
|
||||
func migrateV13(ctx Context, db *gorm.DB, backend string) error {
|
||||
return ctx.EnsureDefaultWAFRuleGroup(db)
|
||||
}
|
||||
|
||||
func validateV13(ctx Context, db *gorm.DB, backend string) error {
|
||||
return ctx.ValidateDatabaseSchemaVersion(db, backend, 13)
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
// v14 升级内容:为 WAF 规则组增加 PoW 策略字段。
|
||||
// 背景说明:PoW 能力从站点路由侧沉淀到 WAF 规则组中,便于统一按规则组管理人机挑战策略。
|
||||
package migrate
|
||||
|
||||
import "gorm.io/gorm"
|
||||
|
||||
func init() {
|
||||
Register(V14())
|
||||
}
|
||||
|
||||
func V14() Migration {
|
||||
return Migration{
|
||||
FromVersion: 13,
|
||||
ToVersion: 14,
|
||||
Migrate: migrateV14,
|
||||
Validate: validateV14,
|
||||
}
|
||||
}
|
||||
|
||||
func migrateV14(ctx Context, db *gorm.DB, backend string) error {
|
||||
return ctx.EnsureDefaultWAFRuleGroup(db)
|
||||
}
|
||||
|
||||
func validateV14(ctx Context, db *gorm.DB, backend string) error {
|
||||
return ctx.ValidateDatabaseSchemaVersion(db, backend, 14)
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
// v15 升级内容:为 nodes 增加 ip_manual_override 字段。
|
||||
// 背景说明:管理端手动指定节点 IP 后,Agent 心跳不应继续覆盖该值,因此需要在节点表中记录 IP 是否由管理端锁定。
|
||||
package migrate
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type nodeV15 struct {
|
||||
IPManualOverride bool `gorm:"column:ip_manual_override;not null;default:false"`
|
||||
}
|
||||
|
||||
func init() {
|
||||
Register(V15())
|
||||
}
|
||||
|
||||
func V15() Migration {
|
||||
return Migration{
|
||||
FromVersion: 14,
|
||||
ToVersion: 15,
|
||||
Migrate: migrateV15,
|
||||
Validate: validateV15,
|
||||
}
|
||||
}
|
||||
|
||||
func (nodeV15) TableName() string {
|
||||
return "nodes"
|
||||
}
|
||||
|
||||
func migrateV15(ctx Context, db *gorm.DB, backend string) error {
|
||||
if db == nil {
|
||||
return fmt.Errorf("database handle is nil")
|
||||
}
|
||||
if !db.Migrator().HasColumn(&nodeV15{}, "ip_manual_override") {
|
||||
if err := db.Migrator().AddColumn(&nodeV15{}, "IPManualOverride"); err != nil {
|
||||
return fmt.Errorf("add nodes.ip_manual_override: %w", err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateV15(ctx Context, db *gorm.DB, backend string) error {
|
||||
if err := ctx.ValidateDatabaseSchemaVersion(db, backend, 14); err != nil {
|
||||
return err
|
||||
}
|
||||
if db == nil || !db.Migrator().HasColumn(&nodeV15{}, "ip_manual_override") {
|
||||
return fmt.Errorf("column nodes.ip_manual_override is missing")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,185 @@
|
||||
// v16 is the first database migration after the V15 formal release baseline.
|
||||
// It folds the previously drafted v16-v21 schema work into a single official
|
||||
// upgrade: tunnel-relay fields, WAF IP groups, current node identity/version
|
||||
// columns, and split node observation tables. The migration also backfills
|
||||
// legacy node columns and removes obsolete pre-release tunnel metadata when
|
||||
// present, so V15 deployments can upgrade directly to the new formal schema.
|
||||
package migrate
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func (nodeV16) TableName() string {
|
||||
return "nodes"
|
||||
}
|
||||
|
||||
func (tunnelV16) TableName() string {
|
||||
return "tunnels"
|
||||
}
|
||||
|
||||
func (proxyRouteV16) TableName() string {
|
||||
return "proxy_routes"
|
||||
}
|
||||
|
||||
type nodeV16 struct{}
|
||||
|
||||
type tunnelV16 struct{}
|
||||
|
||||
type proxyRouteV16 struct{}
|
||||
|
||||
type wafIPGroupV16 struct{}
|
||||
|
||||
type wafRuleGroupV16 struct{}
|
||||
|
||||
func (wafIPGroupV16) TableName() string {
|
||||
return "waf_ip_groups"
|
||||
}
|
||||
|
||||
func (wafRuleGroupV16) TableName() string {
|
||||
return "waf_rule_groups"
|
||||
}
|
||||
|
||||
func init() {
|
||||
Register(V16())
|
||||
}
|
||||
|
||||
func V16() Migration {
|
||||
return Migration{
|
||||
FromVersion: 15,
|
||||
ToVersion: 16,
|
||||
Migrate: migrateV16,
|
||||
Validate: validateV16,
|
||||
}
|
||||
}
|
||||
|
||||
func migrateV16(ctx Context, db *gorm.DB, backend string) error {
|
||||
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
migrator := db.Migrator()
|
||||
if migrator.HasColumn(&nodeV16{}, "agent_token") {
|
||||
if err := db.Exec(`UPDATE nodes SET access_token = agent_token WHERE access_token IS NULL OR access_token = ''`).Error; err != nil {
|
||||
return fmt.Errorf("backfill nodes.access_token from agent_token: %w", err)
|
||||
}
|
||||
}
|
||||
if migrator.HasColumn(&nodeV16{}, "agent_version") {
|
||||
if err := db.Exec(`UPDATE nodes SET version = agent_version WHERE version = '' OR version IS NULL`).Error; err != nil {
|
||||
return fmt.Errorf("backfill nodes.version from agent_version: %w", err)
|
||||
}
|
||||
}
|
||||
if migrator.HasColumn(&nodeV16{}, "nginx_version") {
|
||||
if err := db.Exec(`UPDATE nodes SET ext_version = nginx_version WHERE ext_version IS NULL OR ext_version = ''`).Error; err != nil {
|
||||
return fmt.Errorf("backfill nodes.ext_version from nginx_version: %w", err)
|
||||
}
|
||||
}
|
||||
if err := ctx.DropLegacyNodeColumns(db, backend); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := db.Exec("UPDATE nodes SET node_type = 'edge_node' WHERE node_type = '' OR node_type IS NULL").Error; err != nil {
|
||||
return fmt.Errorf("backfill nodes.node_type: %w", err)
|
||||
}
|
||||
if err := db.Exec("UPDATE proxy_routes SET upstream_type = 'direct' WHERE upstream_type = '' OR upstream_type IS NULL").Error; err != nil {
|
||||
return fmt.Errorf("backfill proxy_routes.upstream_type: %w", err)
|
||||
}
|
||||
|
||||
if migrator.HasColumn(&proxyRouteV16{}, "tunnel_id") {
|
||||
if err := db.Model(&proxyRouteV16{}).Where("upstream_type = ?", "tunnel").Update("upstream_type", "direct").Error; err != nil {
|
||||
return fmt.Errorf("reset pre-release tunnel proxy routes: %w", err)
|
||||
}
|
||||
// Drop the legacy index idx_proxy_routes_tunnel_id if it exists, to avoid errors on dropping the tunnel_id column (especially on SQLite).
|
||||
if migrator.HasIndex(&proxyRouteV16{}, "idx_proxy_routes_tunnel_id") {
|
||||
if err := migrator.DropIndex(&proxyRouteV16{}, "idx_proxy_routes_tunnel_id"); err != nil {
|
||||
return fmt.Errorf("drop index idx_proxy_routes_tunnel_id failed: %w", err)
|
||||
}
|
||||
}
|
||||
if err := migrator.DropColumn(&proxyRouteV16{}, "tunnel_id"); err != nil {
|
||||
return fmt.Errorf("drop pre-release proxy_routes.tunnel_id: %w", err)
|
||||
}
|
||||
}
|
||||
if migrator.HasTable(&tunnelV16{}) {
|
||||
if err := migrator.DropTable(&tunnelV16{}); err != nil {
|
||||
return fmt.Errorf("drop pre-release tunnels table: %w", err)
|
||||
}
|
||||
slog.Info("dropped pre-release tunnels table during v16 migration")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateV16(ctx Context, db *gorm.DB, backend string) error {
|
||||
if err := ctx.ValidateDatabaseSchemaVersion(db, backend, 15); err != nil {
|
||||
return err
|
||||
}
|
||||
if db == nil {
|
||||
return fmt.Errorf("database handle is nil")
|
||||
}
|
||||
|
||||
migrator := db.Migrator()
|
||||
for _, column := range []string{
|
||||
"access_token",
|
||||
"version",
|
||||
"ext_version",
|
||||
"node_type",
|
||||
"relay_bind_port",
|
||||
"relay_vhost_http_port",
|
||||
"relay_auth_token",
|
||||
"relay_agent_access_addr",
|
||||
"relay_client_access_addr",
|
||||
"relay_client_proxy_url",
|
||||
"relay_status",
|
||||
} {
|
||||
if !migrator.HasColumn(&nodeV16{}, column) {
|
||||
return fmt.Errorf("column nodes.%s is missing", column)
|
||||
}
|
||||
}
|
||||
for _, column := range []string{
|
||||
"upstream_type",
|
||||
"tunnel_node_id",
|
||||
"tunnel_target_addr",
|
||||
"tunnel_target_protocol",
|
||||
} {
|
||||
if !migrator.HasColumn(&proxyRouteV16{}, column) {
|
||||
return fmt.Errorf("column proxy_routes.%s is missing", column)
|
||||
}
|
||||
}
|
||||
if migrator.HasColumn(&proxyRouteV16{}, "tunnel_id") {
|
||||
return fmt.Errorf("column proxy_routes.tunnel_id should not exist in v16")
|
||||
}
|
||||
if migrator.HasTable(&tunnelV16{}) {
|
||||
return fmt.Errorf("table tunnels should not exist in v16")
|
||||
}
|
||||
for _, column := range []string{
|
||||
"agent_token",
|
||||
"agent_version",
|
||||
"nginx_version",
|
||||
"relay_version",
|
||||
"relay_frp_version",
|
||||
"relay_frps_connections",
|
||||
"relay_frps_proxy_count",
|
||||
} {
|
||||
if migrator.HasColumn(&nodeV16{}, column) {
|
||||
return fmt.Errorf("column nodes.%s should not exist in v16", column)
|
||||
}
|
||||
}
|
||||
if !migrator.HasTable(&wafIPGroupV16{}) {
|
||||
return fmt.Errorf("table waf_ip_groups is missing")
|
||||
}
|
||||
for _, column := range []string{
|
||||
"ip_whitelist_groups",
|
||||
"ip_blacklist_groups",
|
||||
} {
|
||||
if !migrator.HasColumn(&wafRuleGroupV16{}, column) {
|
||||
return fmt.Errorf("column waf_rule_groups.%s is missing", column)
|
||||
}
|
||||
}
|
||||
if !migrator.HasColumn(&wafIPGroupV16{}, "ext_ips") {
|
||||
return fmt.Errorf("column waf_ip_groups.ext_ips is missing")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package migrate
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type nodeV17 struct{}
|
||||
|
||||
func (nodeV17) TableName() string {
|
||||
return "nodes"
|
||||
}
|
||||
|
||||
func init() {
|
||||
Register(V17())
|
||||
}
|
||||
|
||||
func V17() Migration {
|
||||
return Migration{
|
||||
FromVersion: 16,
|
||||
ToVersion: 17,
|
||||
Migrate: migrateV17,
|
||||
Validate: validateV17,
|
||||
}
|
||||
}
|
||||
|
||||
func migrateV17(ctx Context, db *gorm.DB, backend string) error {
|
||||
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateV17(ctx Context, db *gorm.DB, backend string) error {
|
||||
if err := ctx.ValidateDatabaseSchemaVersion(db, backend, 16); err != nil {
|
||||
return err
|
||||
}
|
||||
if db == nil {
|
||||
return fmt.Errorf("database handle is nil")
|
||||
}
|
||||
|
||||
migrator := db.Migrator()
|
||||
if !migrator.HasColumn(&nodeV17{}, "relay_web_server_enabled") {
|
||||
return fmt.Errorf("column nodes.relay_web_server_enabled is missing")
|
||||
}
|
||||
|
||||
// Validate columns on a sharded partition table
|
||||
for _, shard := range []string{"node_observation_frps_00"} {
|
||||
for _, column := range []string{"frps_client_count", "frps_proxies"} {
|
||||
if !migrator.HasColumn(shard, column) {
|
||||
return fmt.Errorf("column %s.%s is missing", shard, column)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
// v8 升级内容:为 proxy_routes 增加域名级证书绑定字段 domain_cert_ids,并回填已有站点的证书映射。
|
||||
// 背景说明:v1-v7 已作为历史初始基线合并;v8 是当前保留逐版本升级链的起点,用于把早期站点级证书列表扩展为每个域名可独立绑定证书。
|
||||
package migrate
|
||||
|
||||
import "gorm.io/gorm"
|
||||
|
||||
func init() {
|
||||
Register(V8())
|
||||
}
|
||||
|
||||
func V8() Migration {
|
||||
return Migration{
|
||||
FromVersion: 7,
|
||||
ToVersion: 8,
|
||||
Migrate: migrateV8,
|
||||
Validate: validateV8,
|
||||
}
|
||||
}
|
||||
|
||||
func migrateV8(ctx Context, db *gorm.DB, backend string) error {
|
||||
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ctx.BackfillOriginsFromProxyRoutes(db); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ctx.BackfillProxyRouteSiteFields(db); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ctx.EnsureProxyRouteSiteNameUniqueIndex(db); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ctx.BackfillProxyRouteCertificateFields(db); err != nil {
|
||||
return err
|
||||
}
|
||||
return ctx.BackfillProxyRouteDomainCertificateFields(db)
|
||||
}
|
||||
|
||||
func validateV8(ctx Context, db *gorm.DB, backend string) error {
|
||||
return ctx.ValidateDatabaseSchemaVersion(db, backend, 8)
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
// v9 升级内容:为 proxy_routes 增加 PoW 防护配置字段。
|
||||
// 背景说明:反向代理站点需要支持 Proof-of-Work 抗机器人能力,因此在路由配置中持久化 PoW 开关与策略,并沿用 v8 的证书与站点字段回填。
|
||||
package migrate
|
||||
|
||||
import "gorm.io/gorm"
|
||||
|
||||
func init() {
|
||||
Register(V9())
|
||||
}
|
||||
|
||||
func V9() Migration {
|
||||
return Migration{
|
||||
FromVersion: 8,
|
||||
ToVersion: 9,
|
||||
Migrate: migrateV9,
|
||||
Validate: validateV9,
|
||||
}
|
||||
}
|
||||
|
||||
func migrateV9(ctx Context, db *gorm.DB, backend string) error {
|
||||
if err := migrateV8(ctx, db, backend); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateV9(ctx Context, db *gorm.DB, backend string) error {
|
||||
return ctx.ValidateDatabaseSchemaVersion(db, backend, 9)
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,91 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type Node struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
NodeID string `json:"node_id" gorm:"uniqueIndex;size:64;not null"`
|
||||
Name string `json:"name" gorm:"size:128;not null"`
|
||||
IP string `json:"ip" gorm:"size:64;not null"`
|
||||
IPManualOverride bool `json:"ip_manual_override" gorm:"not null;default:false"`
|
||||
GeoName string `json:"geo_name" gorm:"size:128"`
|
||||
GeoLatitude *float64 `json:"geo_latitude"`
|
||||
GeoLongitude *float64 `json:"geo_longitude"`
|
||||
GeoManualOverride bool `json:"geo_manual_override" gorm:"not null;default:false"`
|
||||
AccessToken string `json:"-" gorm:"column:access_token;size:128;index"`
|
||||
AutoUpdateEnabled bool `json:"auto_update_enabled" gorm:"not null;default:false"`
|
||||
UpdateRequested bool `json:"update_requested" gorm:"not null;default:false"`
|
||||
UpdateChannel string `json:"update_channel" gorm:"size:16;not null;default:'stable'"`
|
||||
UpdateTag string `json:"update_tag" gorm:"size:64"`
|
||||
RestartOpenrestyRequested bool `json:"restart_openresty_requested" gorm:"not null;default:false"`
|
||||
Version string `json:"version" gorm:"size:64;not null;default:''"`
|
||||
ExtVersion string `json:"ext_version" gorm:"size:64"`
|
||||
OpenrestyStatus string `json:"openresty_status" gorm:"size:16;not null;default:'unknown'"`
|
||||
OpenrestyMessage string `json:"openresty_message" gorm:"type:text"`
|
||||
Status string `json:"status" gorm:"size:16;not null;default:'offline'"`
|
||||
CurrentVersion string `json:"current_version" gorm:"size:32"`
|
||||
LastSeenAt time.Time `json:"last_seen_at"`
|
||||
LastError string `json:"last_error" gorm:"type:text"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
// Node type: edge_node (default) | tunnel_relay | tunnel_client
|
||||
NodeType string `json:"node_type" gorm:"size:32;not null;default:'edge_node'"`
|
||||
// TunnelRelay specific fields
|
||||
RelayBindPort int `json:"relay_bind_port" gorm:"not null;default:0"`
|
||||
RelayVhostHTTPPort int `json:"relay_vhost_http_port" gorm:"not null;default:0"`
|
||||
RelayAuthToken string `json:"-" gorm:"size:128"`
|
||||
RelayAgentAccessAddr string `json:"relay_agent_access_addr" gorm:"size:255"`
|
||||
RelayClientAccessAddr string `json:"relay_client_access_addr" gorm:"size:255"`
|
||||
RelayClientProxyURL string `json:"relay_client_proxy_url" gorm:"size:512"`
|
||||
CapabilitiesJSON string `json:"capabilities_json" gorm:"type:text;not null;default:'[]'"`
|
||||
RelayStatus string `json:"relay_status" gorm:"size:16;not null;default:'unknown'"`
|
||||
RelayWebServerEnabled bool `json:"relay_web_server_enabled" gorm:"not null;default:false"`
|
||||
}
|
||||
|
||||
func ListNodes() (nodes []*Node, err error) {
|
||||
err = DB.Order("id desc").Find(&nodes).Error
|
||||
return nodes, err
|
||||
}
|
||||
|
||||
func ListNodesByNodeIDs(nodeIDs []string) (nodes []*Node, err error) {
|
||||
if len(nodeIDs) == 0 {
|
||||
return []*Node{}, nil
|
||||
}
|
||||
err = DB.Where("node_id IN ?", nodeIDs).Find(&nodes).Error
|
||||
return nodes, err
|
||||
}
|
||||
|
||||
func GetNodeByNodeID(nodeID string) (*Node, error) {
|
||||
node := &Node{}
|
||||
err := DB.Where("node_id = ?", nodeID).First(node).Error
|
||||
return node, err
|
||||
}
|
||||
|
||||
func GetNodeByID(id uint) (*Node, error) {
|
||||
node := &Node{}
|
||||
err := DB.First(node, id).Error
|
||||
return node, err
|
||||
}
|
||||
|
||||
func GetNodeByAccessToken(token string) (*Node, error) {
|
||||
node := &Node{}
|
||||
err := DB.Where("access_token = ?", token).First(node).Error
|
||||
return node, err
|
||||
}
|
||||
|
||||
func (node *Node) Insert() error {
|
||||
return DB.Create(node).Error
|
||||
}
|
||||
|
||||
func (node *Node) Update() error {
|
||||
return DB.Save(node).Error
|
||||
}
|
||||
|
||||
func (node *Node) Delete() error {
|
||||
return DB.Delete(node).Error
|
||||
}
|
||||
|
||||
func ListNodesByType(nodeType string) (nodes []*Node, err error) {
|
||||
err = DB.Where("node_type = ?", nodeType).Order("id desc").Find(&nodes).Error
|
||||
return nodes, err
|
||||
}
|
||||
@@ -0,0 +1,749 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type NodeAccessLog struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
NodeID string `json:"node_id" gorm:"index:,composite:node_logged_at,priority:1;size:64;not null"`
|
||||
LoggedAt time.Time `json:"logged_at" gorm:"index;index:,composite:node_logged_at,priority:2"`
|
||||
RemoteAddr string `json:"remote_addr" gorm:"index;size:128"`
|
||||
Region string `json:"region" gorm:"size:128"`
|
||||
Host string `json:"host" gorm:"index;size:255"`
|
||||
Path string `json:"path" gorm:"size:2048"`
|
||||
StatusCode int `json:"status_code" gorm:"index"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
type NodeAccessLogRegionCount struct {
|
||||
Region string `json:"region"`
|
||||
Count int64 `json:"count"`
|
||||
}
|
||||
|
||||
type NodeAccessLogQuery struct {
|
||||
NodeID string
|
||||
RemoteAddr string
|
||||
Host string
|
||||
Path string
|
||||
Since time.Time
|
||||
Until time.Time
|
||||
Page int
|
||||
PageSize int
|
||||
SortBy string
|
||||
SortOrder string
|
||||
}
|
||||
|
||||
type NodeAccessLogBucketQuery struct {
|
||||
NodeID string
|
||||
RemoteAddr string
|
||||
Host string
|
||||
Path string
|
||||
Since time.Time
|
||||
Page int
|
||||
PageSize int
|
||||
SortBy string
|
||||
SortOrder string
|
||||
FoldMinutes int
|
||||
}
|
||||
|
||||
type NodeAccessLogBucketRow struct {
|
||||
BucketEpoch int64 `json:"bucket_epoch"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
UniqueIPCount int64 `json:"unique_ip_count"`
|
||||
UniqueHostCount int64 `json:"unique_host_count"`
|
||||
SuccessCount int64 `json:"success_count"`
|
||||
ClientErrorCount int64 `json:"client_error_count"`
|
||||
ServerErrorCount int64 `json:"server_error_count"`
|
||||
}
|
||||
|
||||
type NodeAccessLogBucketIPQuery struct {
|
||||
NodeID string
|
||||
RemoteAddr string
|
||||
Host string
|
||||
Path string
|
||||
BucketStartedAt time.Time
|
||||
FoldMinutes int
|
||||
Page int
|
||||
PageSize int
|
||||
SortBy string
|
||||
SortOrder string
|
||||
}
|
||||
|
||||
type NodeAccessLogBucketIPRow struct {
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
SuccessCount int64 `json:"success_count"`
|
||||
ClientErrorCount int64 `json:"client_error_count"`
|
||||
ServerErrorCount int64 `json:"server_error_count"`
|
||||
LastSeenEpoch int64 `json:"last_seen_epoch"`
|
||||
}
|
||||
|
||||
type NodeAccessLogIPSummaryQuery struct {
|
||||
NodeID string
|
||||
RemoteAddr string
|
||||
Host string
|
||||
Since time.Time
|
||||
Page int
|
||||
PageSize int
|
||||
SortBy string
|
||||
SortOrder string
|
||||
}
|
||||
|
||||
type NodeAccessLogIPSummaryRow struct {
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
TotalRequests int64 `json:"total_requests"`
|
||||
RecentRequests int64 `json:"recent_requests"`
|
||||
LastSeenEpoch int64 `json:"last_seen_epoch"`
|
||||
}
|
||||
|
||||
type NodeAccessLogIPTrendQuery struct {
|
||||
NodeID string
|
||||
RemoteAddr string
|
||||
Host string
|
||||
Since time.Time
|
||||
BucketMinutes int
|
||||
}
|
||||
|
||||
type NodeAccessLogTrendPointRow struct {
|
||||
BucketEpoch int64 `json:"bucket_epoch"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
}
|
||||
|
||||
func (log *NodeAccessLog) BeforeCreate(*gorm.DB) error {
|
||||
return assignObservabilityID(&log.ID)
|
||||
}
|
||||
|
||||
func ListNodeAccessLogs(query NodeAccessLogQuery) (logs []*NodeAccessLog, err error) {
|
||||
all, err := listNodeAccessLogsAcrossShards(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
start, end := paginateBounds(len(all), query.Page, query.PageSize)
|
||||
if start >= len(all) {
|
||||
return []*NodeAccessLog{}, nil
|
||||
}
|
||||
return all[start:end], nil
|
||||
}
|
||||
|
||||
func ListNodeAccessLogsForWAFIPGroup(query NodeAccessLogQuery) ([]*NodeAccessLog, error) {
|
||||
return listNodeAccessLogsAcrossShards(query)
|
||||
}
|
||||
|
||||
func CountNodeAccessLogs(query NodeAccessLogQuery) (totalRecords int64, totalIPs int64, err error) {
|
||||
all, err := listNodeAccessLogsAcrossShards(query)
|
||||
if err != nil {
|
||||
return 0, 0, err
|
||||
}
|
||||
ips := make(map[string]struct{}, len(all))
|
||||
for _, item := range all {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
trimmed := strings.TrimSpace(item.RemoteAddr)
|
||||
if trimmed != "" {
|
||||
ips[trimmed] = struct{}{}
|
||||
}
|
||||
}
|
||||
return int64(len(all)), int64(len(ips)), nil
|
||||
}
|
||||
|
||||
func ListNodeAccessLogRegionCounts(nodeID string, since time.Time, limit int) (items []*NodeAccessLogRegionCount, err error) {
|
||||
logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{
|
||||
NodeID: nodeID,
|
||||
Since: since,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
counts := make(map[string]int64)
|
||||
for _, item := range logs {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
region := strings.TrimSpace(item.Region)
|
||||
if region == "" {
|
||||
continue
|
||||
}
|
||||
counts[region]++
|
||||
}
|
||||
items = make([]*NodeAccessLogRegionCount, 0, len(counts))
|
||||
for region, count := range counts {
|
||||
items = append(items, &NodeAccessLogRegionCount{
|
||||
Region: region,
|
||||
Count: count,
|
||||
})
|
||||
}
|
||||
sort.Slice(items, func(i int, j int) bool {
|
||||
if items[i].Count == items[j].Count {
|
||||
return items[i].Region < items[j].Region
|
||||
}
|
||||
return items[i].Count > items[j].Count
|
||||
})
|
||||
if limit > 0 && len(items) > limit {
|
||||
items = items[:limit]
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func ListNodeAccessLogBuckets(query NodeAccessLogBucketQuery) (items []*NodeAccessLogBucketRow, err error) {
|
||||
rows, err := buildNodeAccessLogBucketRows(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
start, end := paginateBounds(len(rows), query.Page, query.PageSize)
|
||||
if start >= len(rows) {
|
||||
return []*NodeAccessLogBucketRow{}, nil
|
||||
}
|
||||
return rows[start:end], nil
|
||||
}
|
||||
|
||||
func CountNodeAccessLogBuckets(query NodeAccessLogBucketQuery) (total int64, err error) {
|
||||
rows, err := buildNodeAccessLogBucketRows(query)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return int64(len(rows)), nil
|
||||
}
|
||||
|
||||
func ListNodeAccessLogBucketIPs(query NodeAccessLogBucketIPQuery) (items []*NodeAccessLogBucketIPRow, err error) {
|
||||
rows, err := buildNodeAccessLogBucketIPRows(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
start, end := paginateBounds(len(rows), query.Page, query.PageSize)
|
||||
if start >= len(rows) {
|
||||
return []*NodeAccessLogBucketIPRow{}, nil
|
||||
}
|
||||
return rows[start:end], nil
|
||||
}
|
||||
|
||||
func CountNodeAccessLogBucketIPs(query NodeAccessLogBucketIPQuery) (total int64, err error) {
|
||||
rows, err := buildNodeAccessLogBucketIPRows(query)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return int64(len(rows)), nil
|
||||
}
|
||||
|
||||
func ListNodeAccessLogIPSummaries(query NodeAccessLogIPSummaryQuery, recentSince time.Time) (items []*NodeAccessLogIPSummaryRow, err error) {
|
||||
rows, err := buildNodeAccessLogIPSummaryRows(query, recentSince)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
start, end := paginateBounds(len(rows), query.Page, query.PageSize)
|
||||
if start >= len(rows) {
|
||||
return []*NodeAccessLogIPSummaryRow{}, nil
|
||||
}
|
||||
return rows[start:end], nil
|
||||
}
|
||||
|
||||
func CountNodeAccessLogIPSummaries(query NodeAccessLogIPSummaryQuery) (total int64, err error) {
|
||||
rows, err := buildNodeAccessLogIPSummaryRows(query, time.Time{})
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return int64(len(rows)), nil
|
||||
}
|
||||
|
||||
func ListNodeAccessLogIPTrend(query NodeAccessLogIPTrendQuery) (items []*NodeAccessLogTrendPointRow, err error) {
|
||||
logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{
|
||||
NodeID: query.NodeID,
|
||||
RemoteAddr: query.RemoteAddr,
|
||||
Host: query.Host,
|
||||
Since: query.Since,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
remoteAddr := strings.TrimSpace(query.RemoteAddr)
|
||||
if remoteAddr == "" {
|
||||
return []*NodeAccessLogTrendPointRow{}, nil
|
||||
}
|
||||
buckets := make(map[int64]int64)
|
||||
for _, item := range logs {
|
||||
if item == nil || strings.TrimSpace(item.RemoteAddr) != remoteAddr {
|
||||
continue
|
||||
}
|
||||
bucketEpoch := bucketEpochForTime(item.LoggedAt, query.BucketMinutes)
|
||||
buckets[bucketEpoch]++
|
||||
}
|
||||
items = make([]*NodeAccessLogTrendPointRow, 0, len(buckets))
|
||||
for bucketEpoch, requestCount := range buckets {
|
||||
items = append(items, &NodeAccessLogTrendPointRow{
|
||||
BucketEpoch: bucketEpoch,
|
||||
RequestCount: requestCount,
|
||||
})
|
||||
}
|
||||
sort.Slice(items, func(i int, j int) bool {
|
||||
return items[i].BucketEpoch < items[j].BucketEpoch
|
||||
})
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func DeleteNodeAccessLogsBefore(before time.Time) (deleted int64, err error) {
|
||||
return deleteAcrossShards(DB, "node_access_logs", &NodeAccessLog{}, func(tx *gorm.DB) *gorm.DB {
|
||||
return tx.Where("logged_at < ?", before)
|
||||
})
|
||||
}
|
||||
|
||||
func DeleteAllNodeAccessLogs(db *gorm.DB) (deleted int64, err error) {
|
||||
return deleteAcrossShards(db, "node_access_logs", &NodeAccessLog{}, nil)
|
||||
}
|
||||
|
||||
func NodeAccessLogExists(db *gorm.DB, record *NodeAccessLog) (bool, error) {
|
||||
if record == nil {
|
||||
return false, nil
|
||||
}
|
||||
db = normalizeShardedDB(db)
|
||||
for _, table := range observabilityShardTables("node_access_logs") {
|
||||
var count int64
|
||||
if err := db.Table(table).
|
||||
Where(
|
||||
"node_id = ? AND logged_at = ? AND remote_addr = ? AND host = ? AND path = ? AND status_code = ?",
|
||||
record.NodeID,
|
||||
record.LoggedAt,
|
||||
record.RemoteAddr,
|
||||
record.Host,
|
||||
record.Path,
|
||||
record.StatusCode,
|
||||
).
|
||||
Limit(1).
|
||||
Count(&count).Error; err != nil {
|
||||
return false, err
|
||||
}
|
||||
if count > 0 {
|
||||
return true, nil
|
||||
}
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
func DeleteNodeAccessLogsByNodeBefore(db *gorm.DB, nodeID string, before time.Time) (deleted int64, err error) {
|
||||
return deleteAcrossShards(db, "node_access_logs", &NodeAccessLog{}, func(tx *gorm.DB) *gorm.DB {
|
||||
return tx.Where("node_id = ? AND logged_at < ?", nodeID, before)
|
||||
})
|
||||
}
|
||||
|
||||
func applyNodeAccessLogFilters(db *gorm.DB, query NodeAccessLogQuery) *gorm.DB {
|
||||
if trimmed := strings.TrimSpace(query.NodeID); trimmed != "" {
|
||||
db = db.Where("node_id LIKE ?", "%"+trimmed+"%")
|
||||
}
|
||||
if trimmed := strings.TrimSpace(query.RemoteAddr); trimmed != "" {
|
||||
db = db.Where("remote_addr LIKE ?", "%"+trimmed+"%")
|
||||
}
|
||||
if trimmed := strings.TrimSpace(query.Host); trimmed != "" {
|
||||
db = db.Where("host LIKE ?", "%"+trimmed+"%")
|
||||
}
|
||||
if trimmed := strings.TrimSpace(query.Path); trimmed != "" {
|
||||
db = db.Where("path LIKE ?", "%"+trimmed+"%")
|
||||
}
|
||||
if !query.Since.IsZero() {
|
||||
db = db.Where("logged_at >= ?", query.Since)
|
||||
}
|
||||
if !query.Until.IsZero() {
|
||||
db = db.Where("logged_at < ?", query.Until)
|
||||
}
|
||||
return db
|
||||
}
|
||||
|
||||
func listNodeAccessLogsAcrossShards(query NodeAccessLogQuery) ([]*NodeAccessLog, error) {
|
||||
items, err := queryAcrossShards("node_access_logs", func(tx *gorm.DB) ([]*NodeAccessLog, error) {
|
||||
var shardRows []*NodeAccessLog
|
||||
if err := applyNodeAccessLogFilters(tx, query).Find(&shardRows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return shardRows, nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sortNodeAccessLogs(items, query.SortBy, query.SortOrder)
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func buildNodeAccessLogBucketRows(query NodeAccessLogBucketQuery) ([]*NodeAccessLogBucketRow, error) {
|
||||
logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{
|
||||
NodeID: query.NodeID,
|
||||
RemoteAddr: query.RemoteAddr,
|
||||
Host: query.Host,
|
||||
Path: query.Path,
|
||||
Since: query.Since,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
type bucketAccumulator struct {
|
||||
requestCount int64
|
||||
uniqueIPs map[string]struct{}
|
||||
uniqueHosts map[string]struct{}
|
||||
successCount int64
|
||||
clientErrorCount int64
|
||||
serverErrorCount int64
|
||||
}
|
||||
accumulators := make(map[int64]*bucketAccumulator)
|
||||
for _, item := range logs {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
bucketEpoch := bucketEpochForTime(item.LoggedAt, query.FoldMinutes)
|
||||
accumulator := accumulators[bucketEpoch]
|
||||
if accumulator == nil {
|
||||
accumulator = &bucketAccumulator{
|
||||
uniqueIPs: make(map[string]struct{}),
|
||||
uniqueHosts: make(map[string]struct{}),
|
||||
}
|
||||
accumulators[bucketEpoch] = accumulator
|
||||
}
|
||||
accumulator.requestCount++
|
||||
if trimmed := strings.TrimSpace(item.RemoteAddr); trimmed != "" {
|
||||
accumulator.uniqueIPs[trimmed] = struct{}{}
|
||||
}
|
||||
if trimmed := strings.TrimSpace(item.Host); trimmed != "" {
|
||||
accumulator.uniqueHosts[trimmed] = struct{}{}
|
||||
}
|
||||
switch {
|
||||
case item.StatusCode < 400:
|
||||
accumulator.successCount++
|
||||
case item.StatusCode < 500:
|
||||
accumulator.clientErrorCount++
|
||||
default:
|
||||
accumulator.serverErrorCount++
|
||||
}
|
||||
}
|
||||
rows := make([]*NodeAccessLogBucketRow, 0, len(accumulators))
|
||||
for bucketEpoch, accumulator := range accumulators {
|
||||
rows = append(rows, &NodeAccessLogBucketRow{
|
||||
BucketEpoch: bucketEpoch,
|
||||
RequestCount: accumulator.requestCount,
|
||||
UniqueIPCount: int64(len(accumulator.uniqueIPs)),
|
||||
UniqueHostCount: int64(len(accumulator.uniqueHosts)),
|
||||
SuccessCount: accumulator.successCount,
|
||||
ClientErrorCount: accumulator.clientErrorCount,
|
||||
ServerErrorCount: accumulator.serverErrorCount,
|
||||
})
|
||||
}
|
||||
sortNodeAccessLogBucketRows(rows, query.SortBy, query.SortOrder)
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
func buildNodeAccessLogBucketIPRows(query NodeAccessLogBucketIPQuery) ([]*NodeAccessLogBucketIPRow, error) {
|
||||
if query.BucketStartedAt.IsZero() {
|
||||
return []*NodeAccessLogBucketIPRow{}, nil
|
||||
}
|
||||
foldMinutes := query.FoldMinutes
|
||||
if foldMinutes <= 0 {
|
||||
foldMinutes = 3
|
||||
}
|
||||
bucketStartedAt := query.BucketStartedAt.UTC()
|
||||
logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{
|
||||
NodeID: query.NodeID,
|
||||
RemoteAddr: query.RemoteAddr,
|
||||
Host: query.Host,
|
||||
Path: query.Path,
|
||||
Since: bucketStartedAt,
|
||||
Until: bucketStartedAt.Add(time.Duration(foldMinutes) * time.Minute),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
type accumulator struct {
|
||||
requestCount int64
|
||||
successCount int64
|
||||
clientErrorCount int64
|
||||
serverErrorCount int64
|
||||
lastSeenAt time.Time
|
||||
}
|
||||
accumulators := make(map[string]*accumulator)
|
||||
for _, item := range logs {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
remoteAddr := strings.TrimSpace(item.RemoteAddr)
|
||||
if remoteAddr == "" {
|
||||
continue
|
||||
}
|
||||
acc := accumulators[remoteAddr]
|
||||
if acc == nil {
|
||||
acc = &accumulator{}
|
||||
accumulators[remoteAddr] = acc
|
||||
}
|
||||
acc.requestCount++
|
||||
switch {
|
||||
case item.StatusCode < 400:
|
||||
acc.successCount++
|
||||
case item.StatusCode < 500:
|
||||
acc.clientErrorCount++
|
||||
default:
|
||||
acc.serverErrorCount++
|
||||
}
|
||||
if item.LoggedAt.After(acc.lastSeenAt) {
|
||||
acc.lastSeenAt = item.LoggedAt
|
||||
}
|
||||
}
|
||||
rows := make([]*NodeAccessLogBucketIPRow, 0, len(accumulators))
|
||||
for remoteAddr, acc := range accumulators {
|
||||
rows = append(rows, &NodeAccessLogBucketIPRow{
|
||||
RemoteAddr: remoteAddr,
|
||||
RequestCount: acc.requestCount,
|
||||
SuccessCount: acc.successCount,
|
||||
ClientErrorCount: acc.clientErrorCount,
|
||||
ServerErrorCount: acc.serverErrorCount,
|
||||
LastSeenEpoch: acc.lastSeenAt.Unix(),
|
||||
})
|
||||
}
|
||||
sortNodeAccessLogBucketIPRows(rows, query.SortBy, query.SortOrder)
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
func buildNodeAccessLogIPSummaryRows(query NodeAccessLogIPSummaryQuery, recentSince time.Time) ([]*NodeAccessLogIPSummaryRow, error) {
|
||||
logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{
|
||||
NodeID: query.NodeID,
|
||||
RemoteAddr: query.RemoteAddr,
|
||||
Host: query.Host,
|
||||
Since: query.Since,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
type accumulator struct {
|
||||
totalRequests int64
|
||||
recentRequests int64
|
||||
lastSeenAt time.Time
|
||||
}
|
||||
accumulators := make(map[string]*accumulator)
|
||||
for _, item := range logs {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
remoteAddr := strings.TrimSpace(item.RemoteAddr)
|
||||
if remoteAddr == "" {
|
||||
continue
|
||||
}
|
||||
acc := accumulators[remoteAddr]
|
||||
if acc == nil {
|
||||
acc = &accumulator{}
|
||||
accumulators[remoteAddr] = acc
|
||||
}
|
||||
acc.totalRequests++
|
||||
if !recentSince.IsZero() && !item.LoggedAt.Before(recentSince) {
|
||||
acc.recentRequests++
|
||||
}
|
||||
if item.LoggedAt.After(acc.lastSeenAt) {
|
||||
acc.lastSeenAt = item.LoggedAt
|
||||
}
|
||||
}
|
||||
rows := make([]*NodeAccessLogIPSummaryRow, 0, len(accumulators))
|
||||
for remoteAddr, acc := range accumulators {
|
||||
rows = append(rows, &NodeAccessLogIPSummaryRow{
|
||||
RemoteAddr: remoteAddr,
|
||||
TotalRequests: acc.totalRequests,
|
||||
RecentRequests: acc.recentRequests,
|
||||
LastSeenEpoch: acc.lastSeenAt.Unix(),
|
||||
})
|
||||
}
|
||||
sortNodeAccessLogIPSummaryRows(rows, query.SortBy, query.SortOrder)
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
func sortNodeAccessLogBucketIPRows(items []*NodeAccessLogBucketIPRow, sortBy string, sortOrder string) {
|
||||
desc := normalizeSortOrder(sortOrder) != "asc"
|
||||
sort.Slice(items, func(i int, j int) bool {
|
||||
left := items[i]
|
||||
right := items[j]
|
||||
if left == nil || right == nil {
|
||||
return left != nil
|
||||
}
|
||||
var compare int
|
||||
switch strings.TrimSpace(sortBy) {
|
||||
case "last_seen_at":
|
||||
compare = compareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
|
||||
case "remote_addr":
|
||||
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
|
||||
default:
|
||||
compare = compareInt64(left.RequestCount, right.RequestCount)
|
||||
}
|
||||
if compare == 0 {
|
||||
compare = compareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
|
||||
}
|
||||
if compare == 0 {
|
||||
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
|
||||
}
|
||||
if desc {
|
||||
return compare > 0
|
||||
}
|
||||
return compare < 0
|
||||
})
|
||||
}
|
||||
|
||||
func sortNodeAccessLogs(items []*NodeAccessLog, sortBy string, sortOrder string) {
|
||||
desc := normalizeSortOrder(sortOrder) != "asc"
|
||||
sort.Slice(items, func(i int, j int) bool {
|
||||
left := items[i]
|
||||
right := items[j]
|
||||
if left == nil || right == nil {
|
||||
return left != nil
|
||||
}
|
||||
var compare int
|
||||
switch strings.TrimSpace(sortBy) {
|
||||
case "status_code":
|
||||
compare = compareInt(left.StatusCode, right.StatusCode)
|
||||
case "remote_addr":
|
||||
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
|
||||
case "host":
|
||||
compare = strings.Compare(left.Host, right.Host)
|
||||
case "path":
|
||||
compare = strings.Compare(left.Path, right.Path)
|
||||
default:
|
||||
compare = compareTime(left.LoggedAt, right.LoggedAt)
|
||||
}
|
||||
if compare == 0 {
|
||||
compare = compareTime(left.LoggedAt, right.LoggedAt)
|
||||
}
|
||||
if compare == 0 {
|
||||
compare = compareUint(left.ID, right.ID)
|
||||
}
|
||||
if desc {
|
||||
return compare > 0
|
||||
}
|
||||
return compare < 0
|
||||
})
|
||||
}
|
||||
|
||||
func sortNodeAccessLogBucketRows(items []*NodeAccessLogBucketRow, sortBy string, sortOrder string) {
|
||||
desc := normalizeSortOrder(sortOrder) != "asc"
|
||||
sort.Slice(items, func(i int, j int) bool {
|
||||
left := items[i]
|
||||
right := items[j]
|
||||
if left == nil || right == nil {
|
||||
return left != nil
|
||||
}
|
||||
var compare int
|
||||
switch strings.TrimSpace(sortBy) {
|
||||
case "request_count":
|
||||
compare = compareInt64(left.RequestCount, right.RequestCount)
|
||||
default:
|
||||
compare = compareInt64(left.BucketEpoch, right.BucketEpoch)
|
||||
}
|
||||
if compare == 0 {
|
||||
compare = compareInt64(left.BucketEpoch, right.BucketEpoch)
|
||||
}
|
||||
if desc {
|
||||
return compare > 0
|
||||
}
|
||||
return compare < 0
|
||||
})
|
||||
}
|
||||
|
||||
func sortNodeAccessLogIPSummaryRows(items []*NodeAccessLogIPSummaryRow, sortBy string, sortOrder string) {
|
||||
desc := normalizeSortOrder(sortOrder) != "asc"
|
||||
sort.Slice(items, func(i int, j int) bool {
|
||||
left := items[i]
|
||||
right := items[j]
|
||||
if left == nil || right == nil {
|
||||
return left != nil
|
||||
}
|
||||
var compare int
|
||||
switch strings.TrimSpace(sortBy) {
|
||||
case "recent_requests":
|
||||
compare = compareInt64(left.RecentRequests, right.RecentRequests)
|
||||
case "last_seen_at":
|
||||
compare = compareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
|
||||
case "remote_addr":
|
||||
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
|
||||
default:
|
||||
compare = compareInt64(left.TotalRequests, right.TotalRequests)
|
||||
}
|
||||
if compare == 0 {
|
||||
compare = compareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
|
||||
}
|
||||
if compare == 0 {
|
||||
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
|
||||
}
|
||||
if desc {
|
||||
return compare > 0
|
||||
}
|
||||
return compare < 0
|
||||
})
|
||||
}
|
||||
|
||||
func paginateBounds(total int, page int, pageSize int) (int, int) {
|
||||
if page < 0 {
|
||||
page = 0
|
||||
}
|
||||
if pageSize <= 0 {
|
||||
return 0, total
|
||||
}
|
||||
start := page * pageSize
|
||||
if start > total {
|
||||
start = total
|
||||
}
|
||||
end := start + pageSize
|
||||
if end > total {
|
||||
end = total
|
||||
}
|
||||
return start, end
|
||||
}
|
||||
|
||||
func bucketEpochForTime(value time.Time, bucketMinutes int) int64 {
|
||||
bucketSeconds := int64(bucketMinutes * 60)
|
||||
if bucketSeconds <= 0 {
|
||||
bucketSeconds = 180
|
||||
}
|
||||
return (value.UTC().Unix() / bucketSeconds) * bucketSeconds
|
||||
}
|
||||
|
||||
func compareTime(left time.Time, right time.Time) int {
|
||||
switch {
|
||||
case left.After(right):
|
||||
return 1
|
||||
case left.Before(right):
|
||||
return -1
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func compareInt(left int, right int) int {
|
||||
switch {
|
||||
case left > right:
|
||||
return 1
|
||||
case left < right:
|
||||
return -1
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func compareInt64(left int64, right int64) int {
|
||||
switch {
|
||||
case left > right:
|
||||
return 1
|
||||
case left < right:
|
||||
return -1
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func compareUint(left uint, right uint) int {
|
||||
switch {
|
||||
case left > right:
|
||||
return 1
|
||||
case left < right:
|
||||
return -1
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeSortOrder(sortOrder string) string {
|
||||
if strings.EqualFold(strings.TrimSpace(sortOrder), "asc") {
|
||||
return "asc"
|
||||
}
|
||||
return "desc"
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type NodeHealthEvent struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
|
||||
EventType string `json:"event_type" gorm:"index;size:64;not null"`
|
||||
Severity string `json:"severity" gorm:"size:16;not null"`
|
||||
Status string `json:"status" gorm:"index;size:16;not null"`
|
||||
Message string `json:"message" gorm:"type:text"`
|
||||
FirstTriggeredAt time.Time `json:"first_triggered_at" gorm:"index"`
|
||||
LastTriggeredAt time.Time `json:"last_triggered_at" gorm:"index"`
|
||||
ReportedAt time.Time `json:"reported_at" gorm:"index"`
|
||||
ResolvedAt *time.Time `json:"resolved_at" gorm:"index"`
|
||||
MetadataJSON string `json:"metadata_json" gorm:"type:text"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
func GetActiveNodeHealthEvent(nodeID string, eventType string) (*NodeHealthEvent, error) {
|
||||
event := &NodeHealthEvent{}
|
||||
err := DB.Where("node_id = ? AND event_type = ? AND status = ?", nodeID, eventType, "active").First(event).Error
|
||||
return event, err
|
||||
}
|
||||
|
||||
func ListNodeHealthEvents(nodeID string, activeOnly bool, limit int) (events []*NodeHealthEvent, err error) {
|
||||
query := DB.Where("node_id = ?", nodeID).Order("last_triggered_at desc")
|
||||
if activeOnly {
|
||||
query = query.Where("status = ?", "active")
|
||||
}
|
||||
if limit > 0 {
|
||||
query = query.Limit(limit)
|
||||
}
|
||||
err = query.Find(&events).Error
|
||||
return events, err
|
||||
}
|
||||
|
||||
func ListActiveNodeHealthEvents() (events []*NodeHealthEvent, err error) {
|
||||
err = DB.Where("status = ?", "active").Order("last_triggered_at desc").Find(&events).Error
|
||||
return events, err
|
||||
}
|
||||
|
||||
func DeleteNodeHealthEvents(nodeID string) (deleted int64, err error) {
|
||||
result := DB.Where("node_id = ?", nodeID).Delete(&NodeHealthEvent{})
|
||||
return result.RowsAffected, result.Error
|
||||
}
|
||||
@@ -0,0 +1,107 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
const (
|
||||
PagesDeploymentStatusUploaded = "uploaded"
|
||||
PagesDeploymentStatusActive = "active"
|
||||
)
|
||||
|
||||
type PagesProject struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
Name string `json:"name" gorm:"size:255;not null"`
|
||||
Slug string `json:"slug" gorm:"uniqueIndex;size:128;not null"`
|
||||
Description string `json:"description" gorm:"type:text;not null;default:''"`
|
||||
Enabled bool `json:"enabled" gorm:"not null;default:true"`
|
||||
SPAFallbackEnabled bool `json:"spa_fallback_enabled" gorm:"not null;default:false"`
|
||||
SPAFallbackPath string `json:"spa_fallback_path" gorm:"size:512;not null;default:'/index.html'"`
|
||||
APIProxyEnabled bool `json:"api_proxy_enabled" gorm:"not null;default:false"`
|
||||
APIProxyPath string `json:"api_proxy_path" gorm:"size:255;not null;default:''"`
|
||||
APIProxyPass string `json:"api_proxy_pass" gorm:"size:2048;not null;default:''"`
|
||||
APIProxyRewrite string `json:"api_proxy_rewrite" gorm:"size:255;not null;default:''"`
|
||||
ActiveDeploymentID *uint `json:"active_deployment_id" gorm:"index"`
|
||||
RootDir string `json:"root_dir" gorm:"size:512;not null;default:''"`
|
||||
EntryFile string `json:"entry_file" gorm:"size:512;not null;default:'index.html'"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type PagesDeployment struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
ProjectID uint `json:"project_id" gorm:"not null;index"`
|
||||
DeploymentNumber int `json:"deployment_number" gorm:"not null"`
|
||||
Checksum string `json:"checksum" gorm:"size:64;not null;index"`
|
||||
Status string `json:"status" gorm:"size:32;not null;default:'uploaded';index"`
|
||||
ArtifactPath string `json:"artifact_path" gorm:"size:2048;not null"`
|
||||
FileCount int `json:"file_count" gorm:"not null;default:0"`
|
||||
TotalSize int64 `json:"total_size" gorm:"not null;default:0"`
|
||||
CreatedBy string `json:"created_by" gorm:"size:64;not null;default:''"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
ActivatedAt *time.Time `json:"activated_at"`
|
||||
}
|
||||
|
||||
type PagesDeploymentFile struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
DeploymentID uint `json:"deployment_id" gorm:"not null;index"`
|
||||
Path string `json:"path" gorm:"size:2048;not null"`
|
||||
Size int64 `json:"size" gorm:"not null;default:0"`
|
||||
Checksum string `json:"checksum" gorm:"size:64;not null"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
func ListPagesProjects() (projects []*PagesProject, err error) {
|
||||
err = DB.Order("id desc").Find(&projects).Error
|
||||
return projects, err
|
||||
}
|
||||
|
||||
func GetPagesProjectByID(id uint) (*PagesProject, error) {
|
||||
project := &PagesProject{}
|
||||
err := DB.First(project, id).Error
|
||||
return project, err
|
||||
}
|
||||
|
||||
func GetPagesProjectBySlug(slug string) (*PagesProject, error) {
|
||||
project := &PagesProject{}
|
||||
err := DB.Where("slug = ?", slug).First(project).Error
|
||||
return project, err
|
||||
}
|
||||
|
||||
func ListPagesDeployments(projectID uint) (deployments []*PagesDeployment, err error) {
|
||||
err = DB.Where("project_id = ?", projectID).Order("id desc").Find(&deployments).Error
|
||||
return deployments, err
|
||||
}
|
||||
|
||||
func GetPagesDeploymentByID(id uint) (*PagesDeployment, error) {
|
||||
deployment := &PagesDeployment{}
|
||||
err := DB.First(deployment, id).Error
|
||||
return deployment, err
|
||||
}
|
||||
|
||||
func ListPagesDeploymentFiles(deploymentID uint) (files []*PagesDeploymentFile, err error) {
|
||||
err = DB.Where("deployment_id = ?", deploymentID).Order("path asc").Find(&files).Error
|
||||
return files, err
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type ProxyRoute struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
SiteName string `json:"site_name" gorm:"size:255;not null;default:''"`
|
||||
Domain string `json:"domain" gorm:"uniqueIndex;size:255;not null"`
|
||||
Domains string `json:"domains" gorm:"type:text;not null;default:'[]'"`
|
||||
OriginID *uint `json:"origin_id" gorm:"index"`
|
||||
OriginURL string `json:"origin_url" gorm:"size:2048;not null"`
|
||||
OriginHost string `json:"origin_host" gorm:"size:255"`
|
||||
Upstreams string `json:"upstreams" gorm:"type:text;not null;default:'[]'"`
|
||||
Enabled bool `json:"enabled" gorm:"not null;default:true"`
|
||||
EnableHTTPS bool `json:"enable_https" gorm:"column:enable_https;not null;default:false"`
|
||||
CertID *uint `json:"cert_id"`
|
||||
CertIDs string `json:"cert_ids" gorm:"type:text;not null;default:'[]'"`
|
||||
DomainCertIDs string `json:"domain_cert_ids" gorm:"type:text;not null;default:'[]'"`
|
||||
RedirectHTTP bool `json:"redirect_http" gorm:"not null;default:false"`
|
||||
LimitConnPerServer int `json:"limit_conn_per_server" gorm:"not null;default:0"`
|
||||
LimitConnPerIP int `json:"limit_conn_per_ip" gorm:"not null;default:0"`
|
||||
LimitRate string `json:"limit_rate" gorm:"size:32;not null;default:''"`
|
||||
CacheEnabled bool `json:"cache_enabled" gorm:"not null;default:false"`
|
||||
CachePolicy string `json:"cache_policy" gorm:"size:32;not null;default:''"`
|
||||
CacheRules string `json:"cache_rules" gorm:"type:text;not null;default:'[]'"`
|
||||
CustomHeaders string `json:"custom_headers" gorm:"type:text;not null;default:'[]'"`
|
||||
BasicAuthEnabled bool `json:"basic_auth_enabled" gorm:"not null;default:false"`
|
||||
BasicAuthUsername string `json:"basic_auth_username" gorm:"size:255;not null;default:''"`
|
||||
BasicAuthPassword string `json:"basic_auth_password" gorm:"size:255;not null;default:''"`
|
||||
Remark string `json:"remark" gorm:"size:255"`
|
||||
UpstreamType string `json:"upstream_type" gorm:"size:32;not null;default:'direct'"`
|
||||
TunnelNodeID *uint `json:"tunnel_node_id" gorm:"index"`
|
||||
TunnelTargetAddr string `json:"tunnel_target_addr" gorm:"size:512"`
|
||||
TunnelTargetProtocol string `json:"tunnel_target_protocol" gorm:"size:16"`
|
||||
PagesProjectID *uint `json:"pages_project_id" gorm:"index"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
func ListProxyRoutes() (routes []*ProxyRoute, err error) {
|
||||
err = DB.Order("id desc").Find(&routes).Error
|
||||
return routes, err
|
||||
}
|
||||
|
||||
func GetEnabledProxyRoutes() (routes []*ProxyRoute, err error) {
|
||||
err = DB.Where("enabled = ?", true).Order("site_name asc").Order("domain asc").Find(&routes).Error
|
||||
return routes, err
|
||||
}
|
||||
|
||||
func GetProxyRouteByID(id uint) (*ProxyRoute, error) {
|
||||
route := &ProxyRoute{}
|
||||
err := DB.First(route, id).Error
|
||||
return route, err
|
||||
}
|
||||
|
||||
func ListProxyRoutesByOriginID(originID uint) (routes []*ProxyRoute, err error) {
|
||||
err = DB.Where("origin_id = ?", originID).Order("id desc").Find(&routes).Error
|
||||
return routes, err
|
||||
}
|
||||
|
||||
func (route *ProxyRoute) Insert() error {
|
||||
return DB.Create(route).Error
|
||||
}
|
||||
|
||||
func (route *ProxyRoute) Update() error {
|
||||
return DB.Model(&ProxyRoute{}).Where("id = ?", route.ID).Updates(map[string]any{
|
||||
"site_name": route.SiteName,
|
||||
"domain": route.Domain,
|
||||
"domains": route.Domains,
|
||||
"origin_id": route.OriginID,
|
||||
"origin_url": route.OriginURL,
|
||||
"origin_host": route.OriginHost,
|
||||
"upstreams": route.Upstreams,
|
||||
"enabled": route.Enabled,
|
||||
"enable_https": route.EnableHTTPS,
|
||||
"cert_id": route.CertID,
|
||||
"cert_ids": route.CertIDs,
|
||||
"domain_cert_ids": route.DomainCertIDs,
|
||||
"redirect_http": route.RedirectHTTP,
|
||||
"limit_conn_per_server": route.LimitConnPerServer,
|
||||
"limit_conn_per_ip": route.LimitConnPerIP,
|
||||
"limit_rate": route.LimitRate,
|
||||
"cache_enabled": route.CacheEnabled,
|
||||
"cache_policy": route.CachePolicy,
|
||||
"cache_rules": route.CacheRules,
|
||||
"custom_headers": route.CustomHeaders,
|
||||
"basic_auth_enabled": route.BasicAuthEnabled,
|
||||
"basic_auth_username": route.BasicAuthUsername,
|
||||
"basic_auth_password": route.BasicAuthPassword,
|
||||
"remark": route.Remark,
|
||||
"upstream_type": route.UpstreamType,
|
||||
"tunnel_node_id": route.TunnelNodeID,
|
||||
"tunnel_target_addr": route.TunnelTargetAddr,
|
||||
"tunnel_target_protocol": route.TunnelTargetProtocol,
|
||||
"pages_project_id": route.PagesProjectID,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (route *ProxyRoute) Delete() error {
|
||||
return DB.Delete(route).Error
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
schemagoose "github.com/rain-kl/openflare/openflare-server/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)
|
||||
}
|
||||
@@ -0,0 +1,208 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/bwmarrin/snowflake"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/sharding"
|
||||
)
|
||||
|
||||
const observabilityShardCount = 10
|
||||
|
||||
var (
|
||||
observabilityIDNode *snowflake.Node
|
||||
observabilityIDNodeErr error
|
||||
observabilityIDNodeOnce sync.Once
|
||||
)
|
||||
|
||||
func registerSharding(db *gorm.DB, backend string) error {
|
||||
if db == nil {
|
||||
return nil
|
||||
}
|
||||
_ = backend
|
||||
if err := db.Use(sharding.Register(sharding.Config{
|
||||
ShardingKey: "id",
|
||||
NumberOfShards: observabilityShardCount,
|
||||
ShardingAlgorithm: func(value any) (string, error) {
|
||||
return observabilityShardSuffixForValue(value)
|
||||
},
|
||||
ShardingAlgorithmByPrimaryKey: func(id int64) string {
|
||||
return observabilityShardSuffixForInt64(id)
|
||||
},
|
||||
PrimaryKeyGenerator: sharding.PKCustom,
|
||||
PrimaryKeyGeneratorFn: func(tableIdx int64) int64 {
|
||||
return 0
|
||||
},
|
||||
}, shardedObservabilityTables()...)); err != nil {
|
||||
return fmt.Errorf("register observability sharding failed: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func shardedObservabilityTables() []any {
|
||||
return []any{
|
||||
&NodeMetricSnapshot{},
|
||||
&NodeRequestReport{},
|
||||
&NodeAccessLog{},
|
||||
&NodeObservationOpenresty{},
|
||||
&NodeObservationFrps{},
|
||||
&NodeObservationFrpc{},
|
||||
}
|
||||
}
|
||||
|
||||
func shardedObservabilityBaseTables() []string {
|
||||
return []string{
|
||||
"node_metric_snapshots",
|
||||
"node_request_reports",
|
||||
"node_access_logs",
|
||||
"node_observation_openresties",
|
||||
"node_observation_frps",
|
||||
"node_observation_frpcs",
|
||||
}
|
||||
}
|
||||
|
||||
func isShardedObservabilityTable(tableName string) bool {
|
||||
switch strings.TrimSpace(tableName) {
|
||||
case "node_metric_snapshots", "node_request_reports", "node_access_logs", "node_observation_openresties", "node_observation_frps", "node_observation_frpcs":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func observabilityShardTables(baseTable string) []string {
|
||||
tables := make([]string, 0, observabilityShardCount)
|
||||
for _, suffix := range observabilityShardSuffixes() {
|
||||
tables = append(tables, baseTable+suffix)
|
||||
}
|
||||
return tables
|
||||
}
|
||||
|
||||
func observabilityShardSuffixes() []string {
|
||||
suffixes := make([]string, 0, observabilityShardCount)
|
||||
for index := 0; index < observabilityShardCount; index++ {
|
||||
suffixes = append(suffixes, fmt.Sprintf("_%02d", index))
|
||||
}
|
||||
return suffixes
|
||||
}
|
||||
|
||||
func observabilityShardSuffixForID(id uint) string {
|
||||
return fmt.Sprintf("_%02d", uint64(id)%uint64(observabilityShardCount))
|
||||
}
|
||||
|
||||
func observabilityShardSuffixForInt64(id int64) string {
|
||||
if id < 0 {
|
||||
id = -id
|
||||
}
|
||||
return fmt.Sprintf("_%02d", uint64(id)%uint64(observabilityShardCount))
|
||||
}
|
||||
|
||||
func observabilityShardSuffixForValue(value any) (string, error) {
|
||||
switch typed := value.(type) {
|
||||
case int:
|
||||
return observabilityShardSuffixForInt64(int64(typed)), nil
|
||||
case int8:
|
||||
return observabilityShardSuffixForInt64(int64(typed)), nil
|
||||
case int16:
|
||||
return observabilityShardSuffixForInt64(int64(typed)), nil
|
||||
case int32:
|
||||
return observabilityShardSuffixForInt64(int64(typed)), nil
|
||||
case int64:
|
||||
return observabilityShardSuffixForInt64(typed), nil
|
||||
case uint:
|
||||
return observabilityShardSuffixForID(typed), nil
|
||||
case uint8:
|
||||
return observabilityShardSuffixForID(uint(typed)), nil
|
||||
case uint16:
|
||||
return observabilityShardSuffixForID(uint(typed)), nil
|
||||
case uint32:
|
||||
return observabilityShardSuffixForID(uint(typed)), nil
|
||||
case uint64:
|
||||
return fmt.Sprintf("_%02d", typed%uint64(observabilityShardCount)), nil
|
||||
case string:
|
||||
id, err := strconv.ParseUint(strings.TrimSpace(typed), 10, 64)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("invalid sharding id %q", typed)
|
||||
}
|
||||
return fmt.Sprintf("_%02d", id%uint64(observabilityShardCount)), nil
|
||||
default:
|
||||
return "", fmt.Errorf("unsupported observability sharding value type %T", value)
|
||||
}
|
||||
}
|
||||
|
||||
func legacyObservabilityShardTableName(tableName string) string {
|
||||
return tableName + "_legacy_v2_to_v3"
|
||||
}
|
||||
|
||||
func normalizeShardedDB(db *gorm.DB) *gorm.DB {
|
||||
if db != nil {
|
||||
return db
|
||||
}
|
||||
return DB
|
||||
}
|
||||
|
||||
func nextObservabilityID() (uint, error) {
|
||||
observabilityIDNodeOnce.Do(func() {
|
||||
observabilityIDNode, observabilityIDNodeErr = snowflake.NewNode(0)
|
||||
})
|
||||
if observabilityIDNodeErr != nil {
|
||||
return 0, observabilityIDNodeErr
|
||||
}
|
||||
id := observabilityIDNode.Generate().Int64()
|
||||
if id <= 0 {
|
||||
return 0, fmt.Errorf("generated invalid observability id %d", id)
|
||||
}
|
||||
return uint(id), nil
|
||||
}
|
||||
|
||||
func assignObservabilityID(id *uint) error {
|
||||
if id == nil || *id != 0 {
|
||||
return nil
|
||||
}
|
||||
generated, err := nextObservabilityID()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
*id = generated
|
||||
return nil
|
||||
}
|
||||
|
||||
func queryAcrossShards[T any](baseTable string, query func(tx *gorm.DB) ([]T, error)) ([]T, error) {
|
||||
return queryAcrossShardsWithDB(DB, baseTable, query)
|
||||
}
|
||||
|
||||
func queryAcrossShardsWithDB[T any](db *gorm.DB, baseTable string, query func(tx *gorm.DB) ([]T, error)) ([]T, error) {
|
||||
items := make([]T, 0)
|
||||
db = normalizeShardedDB(db)
|
||||
for _, table := range observabilityShardTables(baseTable) {
|
||||
rows, err := query(db.Table(table))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items = append(items, rows...)
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func deleteAcrossShards(db *gorm.DB, baseTable string, model any, apply func(tx *gorm.DB) *gorm.DB) (int64, error) {
|
||||
db = normalizeShardedDB(db)
|
||||
var deleted int64
|
||||
for _, table := range observabilityShardTables(baseTable) {
|
||||
tx := db.Table(table)
|
||||
if apply != nil {
|
||||
tx = apply(tx)
|
||||
} else {
|
||||
tx = tx.Session(&gorm.Session{AllowGlobalUpdate: true})
|
||||
}
|
||||
result := tx.Delete(model)
|
||||
if result.Error != nil {
|
||||
return deleted, result.Error
|
||||
}
|
||||
deleted += result.RowsAffected
|
||||
}
|
||||
return deleted, nil
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type TLSCertificate struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
Name string `json:"name" gorm:"uniqueIndex;size:255;not null"`
|
||||
CertPEM string `json:"-" gorm:"type:text;not null"`
|
||||
KeyPEM string `json:"-" gorm:"type:text;not null"`
|
||||
NotBefore time.Time `json:"not_before"`
|
||||
NotAfter time.Time `json:"not_after"`
|
||||
Remark string `json:"remark" gorm:"size:255"`
|
||||
Provider string `json:"provider" gorm:"size:64;default:'upload'"` // upload, acme
|
||||
AcmeAccountID uint `json:"acme_account_id"`
|
||||
DnsAccountID uint `json:"dns_account_id"`
|
||||
KeyAlgorithm string `json:"key_algorithm" gorm:"size:32"`
|
||||
AutoRenew bool `json:"auto_renew"`
|
||||
PrimaryDomain string `json:"primary_domain" gorm:"size:255"`
|
||||
OtherDomains string `json:"other_domains" gorm:"type:text"`
|
||||
DisableCNAME bool `json:"disable_cname"`
|
||||
SkipDNS bool `json:"skip_dns"`
|
||||
DNS1 string `json:"dns1" gorm:"size:128"`
|
||||
DNS2 string `json:"dns2" gorm:"size:128"`
|
||||
ApplyStatus string `json:"apply_status" gorm:"size:64;default:'ready'"`
|
||||
ApplyMessage string `json:"apply_message" gorm:"type:text"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
func ListTLSCertificates() (certificates []*TLSCertificate, err error) {
|
||||
err = DB.Order("id desc").Find(&certificates).Error
|
||||
return certificates, err
|
||||
}
|
||||
|
||||
func GetTLSCertificateByID(id uint) (*TLSCertificate, error) {
|
||||
certificate := &TLSCertificate{}
|
||||
err := DB.First(certificate, id).Error
|
||||
return certificate, err
|
||||
}
|
||||
|
||||
func (certificate *TLSCertificate) Insert() error {
|
||||
return DB.Create(certificate).Error
|
||||
}
|
||||
|
||||
func (certificate *TLSCertificate) Update() error {
|
||||
return DB.Save(certificate).Error
|
||||
}
|
||||
|
||||
func (certificate *TLSCertificate) Delete() error {
|
||||
return DB.Delete(certificate).Error
|
||||
}
|
||||
@@ -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
Reference in New Issue
Block a user