[优化] go 引用调整

This commit is contained in:
ryan
2026-06-06 10:26:20 +08:00
parent ee1110b752
commit 3cfefb4367
552 changed files with 1642 additions and 2185 deletions
+256
View File
@@ -0,0 +1,256 @@
package acme
import (
"crypto"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/rsa"
"crypto/x509"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"time"
"github.com/go-acme/lego/v4/acme"
"github.com/go-acme/lego/v4/certcrypto"
"github.com/go-acme/lego/v4/certificate"
"github.com/go-acme/lego/v4/challenge/dns01"
"github.com/go-acme/lego/v4/lego"
"github.com/go-acme/lego/v4/providers/dns/cloudflare"
"github.com/go-acme/lego/v4/registration"
)
type AcmeUser struct {
Email string
Registration *registration.Resource
key crypto.PrivateKey
}
func (u *AcmeUser) GetEmail() string {
return u.Email
}
func (u *AcmeUser) GetRegistration() *registration.Resource {
return u.Registration
}
func (u *AcmeUser) GetPrivateKey() crypto.PrivateKey {
return u.key
}
type CertificateResult struct {
CertPEM string
KeyPEM string
NotBefore time.Time
NotAfter time.Time
}
func parsePrivateKey(pemData string) (crypto.PrivateKey, error) {
block, _ := pem.Decode([]byte(pemData))
if block == nil {
return nil, errors.New("failed to parse PEM block containing the key")
}
if key, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil {
return key, nil
}
if key, err := x509.ParsePKCS8PrivateKey(block.Bytes); err == nil {
return key, nil
}
if key, err := x509.ParseECPrivateKey(block.Bytes); err == nil {
return key, nil
}
return nil, errors.New("failed to parse private key")
}
func encodePrivateKey(key crypto.PrivateKey) (string, error) {
var pemBlock *pem.Block
switch k := key.(type) {
case *rsa.PrivateKey:
pemBlock = &pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(k)}
case *ecdsa.PrivateKey:
b, err := x509.MarshalECPrivateKey(k)
if err != nil {
return "", err
}
pemBlock = &pem.Block{Type: "EC PRIVATE KEY", Bytes: b}
default:
return "", errors.New("unsupported key type")
}
return string(pem.EncodeToMemory(pemBlock)), nil
}
func GetOrCreateLegoClient(acmeEmail, privateKeyPEM, accountURL string, keyAlgorithm string) (*lego.Client, *AcmeUser, string, string, error) {
var privateKey crypto.PrivateKey
var err error
var newPrivateKeyPEM string
var newAccountURL string
if privateKeyPEM == "" {
privateKey, err = ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
return nil, nil, "", "", err
}
pemStr, err := encodePrivateKey(privateKey)
if err != nil {
return nil, nil, "", "", err
}
newPrivateKeyPEM = pemStr
} else {
privateKey, err = parsePrivateKey(privateKeyPEM)
if err != nil {
return nil, nil, "", "", err
}
}
user := &AcmeUser{
Email: acmeEmail,
key: privateKey,
}
if accountURL != "" {
user.Registration = &registration.Resource{
Body: acme.Account{
Status: "valid",
Contact: []string{"mailto:" + acmeEmail},
},
URI: accountURL,
}
}
config := lego.NewConfig(user)
// Use Let's Encrypt production environment by default
config.CADirURL = lego.LEDirectoryProduction
switch keyAlgorithm {
case "RSA2048":
config.Certificate.KeyType = certcrypto.RSA2048
case "RSA4096":
config.Certificate.KeyType = certcrypto.RSA4096
case "EC256":
config.Certificate.KeyType = certcrypto.EC256
case "EC384":
config.Certificate.KeyType = certcrypto.EC384
default:
config.Certificate.KeyType = certcrypto.RSA2048
}
client, err := lego.NewClient(config)
if err != nil {
return nil, nil, "", "", err
}
if accountURL == "" {
reg, err := client.Registration.Register(registration.RegisterOptions{TermsOfServiceAgreed: true})
if err != nil {
return nil, nil, "", "", err
}
user.Registration = reg
newAccountURL = reg.URI
}
return client, user, newPrivateKeyPEM, newAccountURL, nil
}
func SetupDNSProvider(client *lego.Client, dnsType, dnsAuth string, dns1, dns2 string, disableCNAME, skipDNS bool) error {
var provider challengeProvider
switch dnsType {
case "cloudflare":
var creds map[string]string
if err := json.Unmarshal([]byte(dnsAuth), &creds); err != nil {
return fmt.Errorf("failed to parse cloudflare credentials: %v", err)
}
config := cloudflare.NewDefaultConfig()
config.AuthToken = creds["api_token"]
p, err := cloudflare.NewDNSProviderConfig(config)
if err != nil {
return err
}
provider = p
default:
return fmt.Errorf("unsupported DNS provider: %s", dnsType)
}
var resolvers []string
if dns1 != "" {
resolvers = append(resolvers, dns1+":53")
}
if dns2 != "" {
resolvers = append(resolvers, dns2+":53")
}
var opts []dns01.ChallengeOption
if len(resolvers) > 0 {
opts = append(opts, dns01.AddRecursiveNameservers(resolvers))
}
if disableCNAME {
opts = append(opts, dns01.DisableCompletePropagationRequirement())
}
if skipDNS {
opts = append(opts, dns01.WrapPreCheck(func(domain, fqdn, value string, check dns01.PreCheckFunc) (bool, error) {
time.Sleep(20 * time.Second)
return true, nil
}))
}
return client.Challenge.SetDNS01Provider(provider, opts...)
}
type challengeProvider interface {
Present(domain, token, keyAuth string) error
CleanUp(domain, token, keyAuth string) error
}
func ObtainSSL(
acmeEmail, acmePrivateKeyPEM, acmeURL string,
dnsType, dnsAuth string,
dns1, dns2 string,
disableCNAME, skipDNS bool,
keyAlgorithm string,
domains []string,
) (string, string, *CertificateResult, error) {
client, _, newPrivateKeyPEM, newAccountURL, err := GetOrCreateLegoClient(acmeEmail, acmePrivateKeyPEM, acmeURL, keyAlgorithm)
if err != nil {
return "", "", nil, fmt.Errorf("failed to create ACME client: %w", err)
}
err = SetupDNSProvider(client, dnsType, dnsAuth, dns1, dns2, disableCNAME, skipDNS)
if err != nil {
return newAccountURL, newPrivateKeyPEM, nil, fmt.Errorf("failed to setup DNS provider: %w", err)
}
request := certificate.ObtainRequest{
Domains: domains,
Bundle: true,
}
certificates, err := client.Certificate.Obtain(request)
if err != nil {
return newAccountURL, newPrivateKeyPEM, nil, fmt.Errorf("failed to obtain certificate: %w", err)
}
result := &CertificateResult{
CertPEM: string(certificates.Certificate),
KeyPEM: string(certificates.PrivateKey),
}
// Parse validity dates
certBlock, _ := pem.Decode(certificates.Certificate)
if certBlock != nil {
parsedCert, err := x509.ParseCertificate(certBlock.Bytes)
if err == nil {
result.NotBefore = parsedCert.NotBefore
result.NotAfter = parsedCert.NotAfter
}
}
return newAccountURL, newPrivateKeyPEM, result, nil
}
+40
View File
@@ -0,0 +1,40 @@
package embedfs
import (
"embed"
"io/fs"
"net/http"
"strings"
"github.com/gin-contrib/static"
)
// Credit: https://github.com/gin-contrib/static/issues/19
type fileSystem struct {
http.FileSystem
}
func (e fileSystem) Exists(prefix string, path string) bool {
cleanPath := strings.TrimPrefix(path, prefix)
cleanPath = strings.TrimPrefix(cleanPath, "/")
if cleanPath == "" {
return false
}
_, err := e.Open(cleanPath)
if err != nil {
return false
}
return true
}
func EmbedFolder(fsEmbed embed.FS, targetPath string) static.ServeFileSystem {
efs, err := fs.Sub(fsEmbed, targetPath)
if err != nil {
panic(err)
}
return fileSystem{
FileSystem: http.FS(efs),
}
}
+53
View File
@@ -0,0 +1,53 @@
package utils
import (
"fmt"
"strconv"
)
var sizeKB = 1024
var sizeMB = sizeKB * 1024
var sizeGB = sizeMB * 1024
func Bytes2Size(num int64) string {
numStr := ""
unit := "B"
if num/int64(sizeGB) > 1 {
numStr = fmt.Sprintf("%.2f", float64(num)/float64(sizeGB))
unit = "GB"
} else if num/int64(sizeMB) > 1 {
numStr = fmt.Sprintf("%d", int(float64(num)/float64(sizeMB)))
unit = "MB"
} else if num/int64(sizeKB) > 1 {
numStr = fmt.Sprintf("%d", int(float64(num)/float64(sizeKB)))
unit = "KB"
} else {
numStr = fmt.Sprintf("%d", num)
}
return numStr + " " + unit
}
func Seconds2Time(num int) (time string) {
if num/31104000 > 0 {
time += strconv.Itoa(num/31104000) + " 年 "
num %= 31104000
}
if num/2592000 > 0 {
time += strconv.Itoa(num/2592000) + " 个月 "
num %= 2592000
}
if num/86400 > 0 {
time += strconv.Itoa(num/86400) + " 天 "
num %= 86400
}
if num/3600 > 0 {
time += strconv.Itoa(num/3600) + " 小时 "
num %= 3600
}
if num/60 > 0 {
time += strconv.Itoa(num/60) + " 分钟 "
num %= 60
}
time += strconv.Itoa(num) + " 秒"
return
}
@@ -0,0 +1,25 @@
package geoip
import (
"fmt"
"net"
)
type EmptyProvider struct{}
func (e *EmptyProvider) Name() string {
return "EmptyProvider"
}
func (e *EmptyProvider) Initialize() error {
return nil
}
func (e *EmptyProvider) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
return nil, fmt.Errorf("you are using an empty GeoIP provider, please set a valid provider")
}
func (e *EmptyProvider) UpdateDatabase() error {
return fmt.Errorf("you are using an empty GeoIP provider, please set a valid provider")
}
func (e *EmptyProvider) Close() error {
return nil
}
+242
View File
@@ -0,0 +1,242 @@
package geoip
import (
"fmt"
"log/slog"
"net"
"strings"
"sync"
"time"
"unicode"
ristretto "github.com/dgraph-io/ristretto/v2"
)
var CurrentProvider GeoIPService
var geoCache *providerCache
var providerMutex sync.RWMutex
var providerFactory = newProvider
const (
ProviderDisabled = "disabled"
ProviderMaxMind = "mmdb"
ProviderIPAPI = "ip-api"
ProviderGeoJS = "geojs"
ProviderIPInfo = "ipinfo"
)
type GeoInfo struct {
ISOCode string
Name string
Latitude *float64
Longitude *float64
}
func init() {
CurrentProvider = &EmptyProvider{}
geoCache = newProviderCache(48 * time.Hour)
}
// GeoIPService 接口定义了获取地理位置信息的核心方法。
type GeoIPService interface {
Name() string
GetGeoInfo(ip net.IP) (*GeoInfo, error)
UpdateDatabase() error
Close() error
}
type cachedGeoInfo struct {
info *GeoInfo
expiresAt time.Time
}
type providerCache struct {
items *ristretto.Cache[string, cachedGeoInfo]
duration time.Duration
}
func newProviderCache(duration time.Duration) *providerCache {
items, err := ristretto.NewCache(&ristretto.Config[string, cachedGeoInfo]{
NumCounters: 1e5,
MaxCost: 2e4,
BufferItems: 64,
})
if err != nil {
panic(err)
}
return &providerCache{
items: items,
duration: duration,
}
}
func (c *providerCache) Get(key string) (*GeoInfo, bool) {
entry, ok := c.items.Get(key)
if !ok {
return nil, false
}
if time.Now().After(entry.expiresAt) {
c.items.Del(key)
return nil, false
}
return entry.info, true
}
func (c *providerCache) Set(key string, info *GeoInfo) {
c.items.Set(key, cachedGeoInfo{
info: info,
expiresAt: time.Now().Add(c.duration),
}, 1)
c.items.Wait()
}
func (c *providerCache) Flush() {
c.items.Clear()
}
func GetRegionUnicodeEmoji(isoCode string) string {
if len(isoCode) != 2 {
return ""
}
isoCode = strings.ToUpper(isoCode)
if !unicode.IsLetter(rune(isoCode[0])) || !unicode.IsLetter(rune(isoCode[1])) {
return ""
}
rune1 := rune(0x1F1E6 + (rune(isoCode[0]) - 'A'))
rune2 := rune(0x1F1E6 + (rune(isoCode[1]) - 'A'))
return string(rune1) + string(rune2)
}
func InitGeoIP(provider string) {
providerName := normalizeProvider(provider)
nextProvider, err := providerFactory(providerName)
if err != nil {
slog.Error("initialize GeoIP provider failed", "provider", providerName, "error", err)
nextProvider = &EmptyProvider{}
}
setProvider(nextProvider)
if providerName == ProviderDisabled {
slog.Info("GeoIP provider disabled")
return
}
slog.Info("GeoIP provider configured", "provider", CurrentProvider.Name())
}
func GetGeoInfo(ip net.IP) (*GeoInfo, error) {
if ip == nil {
return nil, fmt.Errorf("IP address cannot be nil")
}
provider := getProvider()
cacheKey := provider.Name() + ":" + ip.String()
if cachedInfo, found := geoCache.Get(cacheKey); found {
return cachedInfo, nil
}
info, err := provider.GetGeoInfo(ip)
if err == nil && info != nil {
geoCache.Set(cacheKey, info)
}
return info, err
}
func LookupGeoInfoWithProvider(providerName string, ip net.IP) (*GeoInfo, error) {
if ip == nil {
return nil, fmt.Errorf("IP address cannot be nil")
}
provider, err := providerFactory(normalizeProvider(providerName))
if err != nil {
return nil, err
}
defer func() {
if closeErr := provider.Close(); closeErr != nil {
slog.Warn("close temporary GeoIP provider failed", "provider", provider.Name(), "error", closeErr)
}
}()
return provider.GetGeoInfo(ip)
}
func UpdateDatabase() error {
err := getProvider().UpdateDatabase()
if err == nil {
geoCache.Flush()
slog.Info("GeoIP cache cleared due to database update.")
}
return err
}
func IsValidProvider(provider string) bool {
switch normalizeProvider(provider) {
case ProviderDisabled, ProviderMaxMind, ProviderIPAPI, ProviderGeoJS, ProviderIPInfo:
return true
default:
return false
}
}
func normalizeProvider(provider string) string {
normalized := strings.TrimSpace(strings.ToLower(provider))
if normalized == "" {
return ProviderDisabled
}
return normalized
}
func newProvider(provider string) (GeoIPService, error) {
switch provider {
case ProviderDisabled:
return &EmptyProvider{}, nil
case ProviderMaxMind:
return NewMaxMindGeoIPService()
case ProviderIPAPI:
return NewIPAPIService()
case ProviderGeoJS:
return NewGeoJSService()
case ProviderIPInfo:
return NewIPInfoService()
default:
return nil, fmt.Errorf("unsupported GeoIP provider %q", provider)
}
}
func setProvider(provider GeoIPService) {
providerMutex.Lock()
previous := CurrentProvider
CurrentProvider = provider
providerMutex.Unlock()
geoCache.Flush()
if previous != nil && previous != provider {
if err := previous.Close(); err != nil {
slog.Warn("close previous GeoIP provider failed", "error", err)
}
}
}
func getProvider() GeoIPService {
providerMutex.RLock()
defer providerMutex.RUnlock()
if CurrentProvider == nil {
return &EmptyProvider{}
}
return CurrentProvider
}
func float64Pointer(value float64) *float64 {
return &value
}
func ProviderFactoryForTest() func(string) (GeoIPService, error) {
return providerFactory
}
func SetProviderFactoryForTest(factory func(string) (GeoIPService, error)) {
if factory == nil {
providerFactory = newProvider
return
}
providerFactory = factory
}
@@ -0,0 +1,99 @@
package geoip
import (
"net"
"testing"
)
type fakeProvider struct {
calls int
}
func (f *fakeProvider) Name() string {
return "fake"
}
func (f *fakeProvider) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
f.calls++
return &GeoInfo{
ISOCode: "CN",
Name: "China",
}, nil
}
func (f *fakeProvider) UpdateDatabase() error {
return nil
}
func (f *fakeProvider) Close() error {
return nil
}
func TestGetGeoInfoCachesByProviderAndIP(t *testing.T) {
originalProvider := CurrentProvider
geoCache.Flush()
fake := &fakeProvider{}
CurrentProvider = fake
defer func() {
CurrentProvider = originalProvider
}()
ip := net.ParseIP("8.8.8.8")
record, err := GetGeoInfo(ip)
if err != nil {
t.Fatalf("expected nil error, got %v", err)
}
if record == nil || record.ISOCode != "CN" {
t.Fatalf("expected cached record, got %#v", record)
}
_, err = GetGeoInfo(ip)
if err != nil {
t.Fatalf("expected nil error on second call, got %v", err)
}
if fake.calls != 1 {
t.Fatalf("expected provider to be called once, got %d", fake.calls)
}
}
func TestUnicodeEmoji(t *testing.T) {
emoji := GetRegionUnicodeEmoji("CN")
if emoji != "🇨🇳" {
t.Errorf("expected emoji for CN, got %s", emoji)
}
}
func TestIsValidProvider(t *testing.T) {
cases := map[string]bool{
"disabled": true,
"mmdb": true,
"ip-api": true,
"geojs": true,
"ipinfo": true,
"unknown": false,
}
for provider, want := range cases {
if got := IsValidProvider(provider); got != want {
t.Fatalf("provider %s validity mismatch: want %v, got %v", provider, want, got)
}
}
}
func TestLookupGeoInfoWithProviderUsesTemporaryProvider(t *testing.T) {
previousFactory := providerFactory
providerFactory = func(provider string) (GeoIPService, error) {
return &fakeProvider{}, nil
}
defer func() {
providerFactory = previousFactory
}()
info, err := LookupGeoInfoWithProvider("ipinfo", net.ParseIP("8.8.8.8"))
if err != nil {
t.Fatalf("expected lookup to succeed, got %v", err)
}
if info == nil || info.ISOCode != "CN" || info.Name != "China" {
t.Fatalf("unexpected geo info: %#v", info)
}
}
+84
View File
@@ -0,0 +1,84 @@
package geoip
import (
"encoding/json"
"fmt"
"net"
"net/http"
"time"
)
// GeoJSService 使用 geojs.io 服务实现 GeoIPService 接口。
type GeoJSService struct {
Client *http.Client
}
// geoJSResponse 定义了 geojs.io 服务返回的 JSON 响应的结构。
// 我们只定义我们需要的字段。
type geoJSResponse struct {
Country string `json:"country"`
CountryCode string `json:"country_code"`
Latitude float64 `json:"latitude,string"`
Longitude float64 `json:"longitude,string"`
// 可以根据需要添加其他字段,例如:
// City string `json:"city"`
// Region string `json:"region"`
}
// NewGeoJSService 创建并返回一个 GeoJSService 的新实例。
func NewGeoJSService() (*GeoJSService, error) {
return &GeoJSService{
Client: &http.Client{
Timeout: 5 * time.Second, // 设置一个合理的超时时间
},
}, nil
}
// Name 返回服务的名称。
func (s *GeoJSService) Name() string {
return "geojs.io"
}
// GetGeoInfo 使用 geojs.io 服务检索给定 IP 地址的地理位置信息。
func (s *GeoJSService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
// GeoJS 的 API 端点
apiURL := fmt.Sprintf("https://get.geojs.io/v1/ip/geo/%s.json", ip.String())
resp, err := s.Client.Get(apiURL)
if err != nil {
return nil, fmt.Errorf("failed to get geo info from geojs.io: %w", err)
}
defer resp.Body.Close()
// 检查响应状态码
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("geojs.io returned non-200 status code: %d", resp.StatusCode)
}
var apiResp geoJSResponse
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
return nil, fmt.Errorf("failed to decode geojs.io response: %w", err)
}
// 检查国家代码是否为空,因为 geojs 对无效/私有IP可能返回200 OK但内容为空
if apiResp.CountryCode == "" {
return nil, fmt.Errorf("geojs.io returned empty geo info for ip: %s", ip.String())
}
return &GeoInfo{
ISOCode: apiResp.CountryCode,
Name: apiResp.Country,
Latitude: float64Pointer(apiResp.Latitude),
Longitude: float64Pointer(apiResp.Longitude),
}, nil
}
// UpdateDatabase 对于 geojs.io 是一个空操作,因为它是一个 Web 服务。
func (s *GeoJSService) UpdateDatabase() error {
return nil
}
// Close 对于 geojs.io 是一个空操作。
func (s *GeoJSService) Close() error {
return nil
}
+86
View File
@@ -0,0 +1,86 @@
package geoip
import (
"encoding/json"
"fmt"
"net"
"net/http"
"time"
)
// IPAPIService 使用 ip-api.com 服务实现 GeoIPService 接口。
type IPAPIService struct {
Client *http.Client
}
// ipAPIResponse 定义了 ip-api.com 服务返回的 JSON 响应的结构。
type ipAPIResponse struct {
Status string `json:"status"`
Message string `json:"message"` // 当 status 为 fail 时出现
Country string `json:"country"`
CountryCode string `json:"countryCode"`
Region string `json:"region"`
RegionName string `json:"regionName"`
City string `json:"city"`
Zip string `json:"zip"`
Lat float64 `json:"lat"`
Lon float64 `json:"lon"`
Timezone string `json:"timezone"`
ISP string `json:"isp"`
Org string `json:"org"`
As string `json:"as"`
Query string `json:"query"`
}
func (s *IPAPIService) Name() string {
return "ip-api.com"
}
// NewIPAPIService 创建并返回一个 IPAPIService 的新实例。
func NewIPAPIService() (*IPAPIService, error) {
return &IPAPIService{
Client: &http.Client{
Timeout: 5 * time.Second, // 设置请求超时
},
}, nil
}
// GetGeoInfo 使用 ip-api.com 服务检索给定 IP 地址的地理位置信息。
func (s *IPAPIService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
// API URL, 使用 fields 参数来仅请求需要的字段
apiURL := fmt.Sprintf("http://ip-api.com/json/%s?fields=status,message,country,countryCode", ip.String())
resp, err := s.Client.Get(apiURL)
if err != nil {
return nil, fmt.Errorf("failed to get geo info from ip-api.com: %w", err)
}
defer resp.Body.Close()
var apiResp ipAPIResponse
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
return nil, fmt.Errorf("failed to decode ip-api.com response: %w", err)
}
if apiResp.Status != "success" {
return nil, fmt.Errorf("ip-api.com returned an error: %s", apiResp.Message)
}
return &GeoInfo{
ISOCode: apiResp.CountryCode,
Name: apiResp.Country,
Latitude: float64Pointer(apiResp.Lat),
Longitude: float64Pointer(apiResp.Lon),
}, nil
}
// UpdateDatabase 对于 ip-api.com 是一个空操作,因为它是一个 Web 服务。
func (s *IPAPIService) UpdateDatabase() error {
// 无需执行任何操作,因为数据由外部服务提供
return nil
}
// Close 对于 ip-api.com 是一个空操作,因为没有需要关闭的持久连接。
func (s *IPAPIService) Close() error {
// 无需执行任何操作
return nil
}
+112
View File
@@ -0,0 +1,112 @@
package geoip
import (
"encoding/json"
"fmt"
"net"
"net/http"
"strconv"
"strings"
"time"
)
// IPInfoService 使用 ipinfo.io 服务实现 GeoIPService 接口。
type IPInfoService struct {
Client *http.Client
// 每天 1000 次请求,限制由 IP 地址的所有人共享。
// APIToken string
}
// ipInfoResponse 定义了 ipinfo.io 服务返回的 JSON 响应的结构,只包含免费额度可用的字段。
type ipInfoResponse struct {
IP string `json:"ip"`
Hostname string `json:"hostname"`
City string `json:"city"`
Region string `json:"region"`
Country string `json:"country"`
CountryCode string `json:"countryCode"` // ipinfo.io 返回 "country" 的 ISO 代码,这里为了与 GeoInfo 保持一致,额外添加一个 CountryCode
Loc string `json:"loc"` // Latitude,Longitude
Org string `json:"org"`
Postal string `json:"postal"`
Timezone string `json:"timezone"`
}
// NewIPInfoService 创建并返回一个 IPInfoService 的新实例。
func NewIPInfoService() (*IPInfoService, error) {
return &IPInfoService{
Client: &http.Client{
Timeout: 5 * time.Second,
},
}, nil
}
// Name 返回服务的名称。
func (s *IPInfoService) Name() string {
return "ipinfo.io"
}
// GetGeoInfo 使用 ipinfo.io 服务检索给定 IP 地址的地理位置信息。
// 免费额度主要提供国家信息。
func (s *IPInfoService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
// IPinfo 免费额度不需要 API token 就可以查询基本的 IP 信息。
// API URL: https://ipinfo.io/json (查询自身IP) 或 https://ipinfo.io/YOUR_IP/json
apiURL := fmt.Sprintf("https://ipinfo.io/%s/json", ip.String())
resp, err := s.Client.Get(apiURL)
if err != nil {
return nil, fmt.Errorf("failed to get geo info from ipinfo.io: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("ipinfo.io returned non-200 status: %d %s", resp.StatusCode, resp.Status)
}
var apiResp ipInfoResponse
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
return nil, fmt.Errorf("failed to decode ipinfo.io response: %w", err)
}
latitude, longitude := parseIPInfoCoordinates(apiResp.Loc)
// IPinfo 的 "country" 字段直接返回 ISO 2-letter code,例如 "US", "CN"
// 我们需要将 "country" 字段作为 ISOCode,并尝试获取其对应的国家名称。
// IPinfo 响应中通常不直接提供完整的国家名称,但我们可以通过 CountryCode 映射。
// 为了简化并符合 GeoInfo 结构,我们直接使用 Country 作为 ISOCode,并尝试从 CountryCode 获取名称。
// 实际上,IPinfo 的 'country' 字段就是 ISO 2-letter code。
// 如果需要完整的国家名称,可能需要一个本地的 ISO 代码到名称的映射。
// 为了与 GetRegionUnicodeEmoji 函数兼容,我们直接使用 country 作为 ISOCode。
return &GeoInfo{
ISOCode: apiResp.Country,
Name: apiResp.Country,
Latitude: latitude,
Longitude: longitude,
}, nil
}
// UpdateDatabase 对于 ipinfo.io 是一个空操作,因为它是一个 Web 服务。
func (s *IPInfoService) UpdateDatabase() error {
// 无需执行任何操作,因为数据由外部服务提供
return nil
}
// Close 对于 ipinfo.io 是一个空操作,因为没有需要关闭的持久连接。
func (s *IPInfoService) Close() error {
// 无需执行任何操作
return nil
}
func parseIPInfoCoordinates(value string) (*float64, *float64) {
parts := strings.Split(strings.TrimSpace(value), ",")
if len(parts) != 2 {
return nil, nil
}
latitudeValue, latErr := strconv.ParseFloat(strings.TrimSpace(parts[0]), 64)
longitudeValue, lonErr := strconv.ParseFloat(strings.TrimSpace(parts[1]), 64)
if latErr != nil || lonErr != nil {
return nil, nil
}
return float64Pointer(latitudeValue), float64Pointer(longitudeValue)
}
@@ -0,0 +1,66 @@
package iputil
import (
"net"
"strings"
)
func NormalizeIP(raw string) string {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return ""
}
ip := net.ParseIP(trimmed)
if ip == nil {
return ""
}
if ipv4 := ip.To4(); ipv4 != nil {
return ipv4.String()
}
return ip.String()
}
func NormalizeRemoteAddr(remoteAddr string) string {
trimmed := strings.TrimSpace(remoteAddr)
if trimmed == "" {
return ""
}
if host, _, err := net.SplitHostPort(trimmed); err == nil {
return NormalizeIP(host)
}
return NormalizeIP(trimmed)
}
func IsPublic(ip net.IP) bool {
if ip == nil {
return false
}
if ipv4 := ip.To4(); ipv4 != nil {
ip = ipv4
}
if !ip.IsGlobalUnicast() || ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsMulticast() || ip.IsUnspecified() {
return false
}
return true
}
func IsPublicString(raw string) bool {
ip := net.ParseIP(strings.TrimSpace(raw))
return IsPublic(ip)
}
func Score(ip net.IP) int {
if ip == nil {
return -1
}
if ipv4 := ip.To4(); ipv4 != nil {
ip = ipv4
}
if !ip.IsGlobalUnicast() || ip.IsLoopback() || ip.IsMulticast() || ip.IsUnspecified() {
return -1
}
if IsPublic(ip) {
return 2
}
return 1
}
@@ -0,0 +1,45 @@
package iputil
import (
"net"
"testing"
)
func TestNormalizeIP(t *testing.T) {
if got := NormalizeIP(" 8.8.8.8 "); got != "8.8.8.8" {
t.Fatalf("unexpected normalized ipv4: %q", got)
}
if got := NormalizeIP("[::1]"); got != "" {
t.Fatalf("expected invalid bracketed host to be rejected, got %q", got)
}
}
func TestNormalizeRemoteAddr(t *testing.T) {
if got := NormalizeRemoteAddr("203.0.113.10:8443"); got != "203.0.113.10" {
t.Fatalf("unexpected remote addr normalization: %q", got)
}
}
func TestIsPublic(t *testing.T) {
if !IsPublic(net.ParseIP("8.8.8.8")) {
t.Fatal("expected public ip to be detected")
}
if IsPublic(net.ParseIP("10.0.0.8")) {
t.Fatal("expected private ip to be rejected")
}
if IsPublic(net.ParseIP("127.0.0.1")) {
t.Fatal("expected loopback ip to be rejected")
}
}
func TestScore(t *testing.T) {
if got := Score(net.ParseIP("8.8.8.8")); got != 2 {
t.Fatalf("unexpected score for public ip: %d", got)
}
if got := Score(net.ParseIP("10.0.0.8")); got != 1 {
t.Fatalf("unexpected score for private ip: %d", got)
}
if got := Score(net.ParseIP("127.0.0.1")); got != -1 {
t.Fatalf("unexpected score for loopback ip: %d", got)
}
}
+171
View File
@@ -0,0 +1,171 @@
package geoip
import (
"fmt"
"io"
"net"
"net/http"
"os"
"path/filepath"
"sync"
"github.com/oschwald/maxminddb-golang"
)
var GeoIpUrl = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb"
var GeoIpFilePath = "./data/GeoLite2-Country.mmdb"
type GeoIpRecord struct {
Country struct {
ISOCode string `maxminddb:"iso_code"`
Names map[string]string `maxminddb:"names"`
} `maxminddb:"country"`
}
type MaxMindGeoIPService struct {
maxMindDBReader *maxminddb.Reader
dbFilePath string
mu sync.RWMutex
}
func (s *MaxMindGeoIPService) Name() string {
return "MaxMind"
}
func NewMaxMindGeoIPService() (*MaxMindGeoIPService, error) {
return NewMaxMindGeoIPServiceWithConfig(GeoIpFilePath, GeoIpUrl)
}
func NewMaxMindGeoIPServiceWithConfig(dbFilePath string, downloadURL string) (*MaxMindGeoIPService, error) {
if dbFilePath == "" {
dbFilePath = GeoIpFilePath
}
if downloadURL == "" {
downloadURL = GeoIpUrl
}
service := &MaxMindGeoIPService{
dbFilePath: dbFilePath,
}
if err := os.MkdirAll(filepath.Dir(service.dbFilePath), os.ModePerm); err != nil {
return nil, fmt.Errorf("failed to create data directory for MaxMind database: %w", err)
}
if _, err := os.Stat(service.dbFilePath); os.IsNotExist(err) {
if err := DownloadMaxMindDatabase(service.dbFilePath, downloadURL); err != nil {
return nil, fmt.Errorf("failed to download initial MaxMind database: %w", err)
}
}
if err := service.initialize(); err != nil {
return nil, fmt.Errorf("failed to initialize MaxMind database: %w", err)
}
return service, nil
}
func (s *MaxMindGeoIPService) initialize() error {
s.mu.Lock()
defer s.mu.Unlock()
if s.maxMindDBReader != nil {
_ = s.maxMindDBReader.Close()
s.maxMindDBReader = nil
}
reader, err := maxminddb.Open(s.dbFilePath)
if err != nil {
return fmt.Errorf("error opening MaxMind database at %s: %w", s.dbFilePath, err)
}
s.maxMindDBReader = reader
return nil
}
func (s *MaxMindGeoIPService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
s.mu.RLock()
defer s.mu.RUnlock()
if s.maxMindDBReader == nil {
return nil, fmt.Errorf("MaxMind database is not initialized or failed to open")
}
if ip == nil {
return nil, fmt.Errorf("IP address cannot be nil")
}
var record GeoIpRecord
if err := s.maxMindDBReader.Lookup(ip, &record); err != nil {
return nil, fmt.Errorf("error looking up IP %s in MaxMind database: %w", ip.String(), err)
}
geoInfo := &GeoInfo{
ISOCode: record.Country.ISOCode,
Name: record.Country.Names["en"],
}
if geoInfo.Name == "" && geoInfo.ISOCode != "" {
geoInfo.Name = geoInfo.ISOCode
}
return geoInfo, nil
}
func (s *MaxMindGeoIPService) UpdateDatabase() error {
if err := DownloadMaxMindDatabase(s.dbFilePath, GeoIpUrl); err != nil {
return err
}
return s.initialize()
}
func DownloadMaxMindDatabase(dbFilePath string, downloadURL string) error {
if dbFilePath == "" {
dbFilePath = GeoIpFilePath
}
if downloadURL == "" {
downloadURL = GeoIpUrl
}
resp, err := http.Get(downloadURL)
if err != nil {
return fmt.Errorf("failed to initiate MaxMind database download: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("failed to download MaxMind database: HTTP status %s", resp.Status)
}
if err := os.MkdirAll(filepath.Dir(dbFilePath), os.ModePerm); err != nil {
return fmt.Errorf("failed to create data directory for MaxMind database update: %w", err)
}
tempPath := dbFilePath + ".download"
out, err := os.Create(tempPath)
if err != nil {
return fmt.Errorf("failed to create MaxMind database file at %s: %w", tempPath, err)
}
defer func() {
_ = out.Close()
}()
if _, err = io.Copy(out, resp.Body); err != nil {
return fmt.Errorf("failed to write MaxMind database file: %w", err)
}
if err = out.Close(); err != nil {
return fmt.Errorf("failed to close MaxMind database file: %w", err)
}
if err = os.Rename(tempPath, dbFilePath); err != nil {
return fmt.Errorf("failed to move MaxMind database file into place: %w", err)
}
return nil
}
func (s *MaxMindGeoIPService) Close() error {
s.mu.Lock()
defer s.mu.Unlock()
if s.maxMindDBReader != nil {
err := s.maxMindDBReader.Close()
s.maxMindDBReader = nil
if err != nil {
return fmt.Errorf("error closing MaxMind database: %w", err)
}
}
return nil
}
+152
View File
@@ -0,0 +1,152 @@
package geoip
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"strings"
"time"
"github.com/rain-kl/openflare/openflare-server/utils/geoip/iputil"
)
const defaultOutboundIPLookupTimeout = 5 * time.Second
// OutboundIPStrategy defines a lookup strategy for the current public egress IP.
type OutboundIPStrategy interface {
Name() string
GetOutboundIP(ctx context.Context) (net.IP, error)
}
// OutboundIPAPIAdapter adapts a third-party HTTP API response into an IP value.
type OutboundIPAPIAdapter interface {
Name() string
Endpoint() string
DecodeIP(io.Reader) (net.IP, error)
}
type HTTPOutboundIPStrategy struct {
Client *http.Client
Adapter OutboundIPAPIAdapter
}
func NewHTTPOutboundIPStrategy(adapter OutboundIPAPIAdapter, client *http.Client) *HTTPOutboundIPStrategy {
if client == nil {
client = &http.Client{Timeout: defaultOutboundIPLookupTimeout}
}
return &HTTPOutboundIPStrategy{
Client: client,
Adapter: adapter,
}
}
func (s *HTTPOutboundIPStrategy) Name() string {
if s == nil || s.Adapter == nil {
return "http-outbound-ip"
}
return s.Adapter.Name()
}
func (s *HTTPOutboundIPStrategy) GetOutboundIP(ctx context.Context) (net.IP, error) {
if s == nil || s.Adapter == nil {
return nil, errors.New("outbound IP adapter is nil")
}
if ctx == nil {
ctx = context.Background()
}
client := s.Client
if client == nil {
client = &http.Client{Timeout: defaultOutboundIPLookupTimeout}
}
request, err := http.NewRequestWithContext(ctx, http.MethodGet, s.Adapter.Endpoint(), nil)
if err != nil {
return nil, fmt.Errorf("%s create request failed: %w", s.Name(), err)
}
response, err := client.Do(request)
if err != nil {
return nil, fmt.Errorf("%s request failed: %w", s.Name(), err)
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
return nil, fmt.Errorf("%s returned non-200 status: %d %s", s.Name(), response.StatusCode, response.Status)
}
ip, err := s.Adapter.DecodeIP(response.Body)
if err != nil {
return nil, fmt.Errorf("%s decode response failed: %w", s.Name(), err)
}
if !iputil.IsPublic(ip) {
return nil, fmt.Errorf("%s returned non-public IP: %s", s.Name(), ip.String())
}
return ip, nil
}
type RealIPCCAdapter struct {
URL string
}
type realIPCCResponse struct {
IP string `json:"ip"`
}
func NewRealIPCCOutboundIPStrategy() *HTTPOutboundIPStrategy {
return NewHTTPOutboundIPStrategy(RealIPCCAdapter{}, nil)
}
func (a RealIPCCAdapter) Name() string {
return "realip.cc"
}
func (a RealIPCCAdapter) Endpoint() string {
if strings.TrimSpace(a.URL) != "" {
return strings.TrimSpace(a.URL)
}
return "https://realip.cc"
}
func (a RealIPCCAdapter) DecodeIP(reader io.Reader) (net.IP, error) {
var payload realIPCCResponse
if err := json.NewDecoder(reader).Decode(&payload); err != nil {
return nil, err
}
ip := net.ParseIP(strings.TrimSpace(payload.IP))
if ip == nil {
return nil, fmt.Errorf("invalid IP %q", payload.IP)
}
if ipv4 := ip.To4(); ipv4 != nil {
return ipv4, nil
}
return ip, nil
}
func DefaultOutboundIPStrategies() []OutboundIPStrategy {
return []OutboundIPStrategy{
NewRealIPCCOutboundIPStrategy(),
}
}
func GetOutboundIP(ctx context.Context, strategies ...OutboundIPStrategy) (net.IP, error) {
if len(strategies) == 0 {
strategies = DefaultOutboundIPStrategies()
}
var errs []error
for _, strategy := range strategies {
if strategy == nil {
continue
}
ip, err := strategy.GetOutboundIP(ctx)
if err == nil && ip != nil {
return ip, nil
}
if err != nil {
errs = append(errs, fmt.Errorf("%s: %w", strategy.Name(), err))
}
}
if len(errs) == 0 {
return nil, errors.New("no outbound IP lookup strategy configured")
}
return nil, errors.Join(errs...)
}
@@ -0,0 +1,82 @@
package geoip
import (
"context"
"errors"
"net"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
type fakeOutboundIPStrategy struct {
name string
ip net.IP
err error
}
func (f fakeOutboundIPStrategy) Name() string {
return f.name
}
func (f fakeOutboundIPStrategy) GetOutboundIP(ctx context.Context) (net.IP, error) {
return f.ip, f.err
}
func TestRealIPCCAdapterDecodeIP(t *testing.T) {
ip, err := RealIPCCAdapter{}.DecodeIP(strings.NewReader(`{"ip":"8.8.8.8","country":"United States"}`))
if err != nil {
t.Fatalf("DecodeIP failed: %v", err)
}
if ip.String() != "8.8.8.8" {
t.Fatalf("unexpected IP: %s", ip.String())
}
}
func TestHTTPOutboundIPStrategyUsesAdapter(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
t.Fatalf("unexpected method: %s", r.Method)
}
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"ip":"8.8.4.4"}`))
}))
defer server.Close()
strategy := NewHTTPOutboundIPStrategy(RealIPCCAdapter{URL: server.URL}, server.Client())
ip, err := strategy.GetOutboundIP(context.Background())
if err != nil {
t.Fatalf("GetOutboundIP failed: %v", err)
}
if ip.String() != "8.8.4.4" {
t.Fatalf("unexpected outbound IP: %s", ip.String())
}
}
func TestGetOutboundIPFallsBackToNextStrategy(t *testing.T) {
ip, err := GetOutboundIP(
context.Background(),
fakeOutboundIPStrategy{name: "first", err: errors.New("temporary failure")},
fakeOutboundIPStrategy{name: "second", ip: net.ParseIP("1.1.1.1")},
)
if err != nil {
t.Fatalf("GetOutboundIP failed: %v", err)
}
if ip.String() != "1.1.1.1" {
t.Fatalf("unexpected outbound IP: %s", ip.String())
}
}
func TestHTTPOutboundIPStrategyRejectsPrivateIP(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"ip":"172.17.0.2"}`))
}))
defer server.Close()
strategy := NewHTTPOutboundIPStrategy(RealIPCCAdapter{URL: server.URL}, server.Client())
if _, err := strategy.GetOutboundIP(context.Background()); err == nil {
t.Fatal("expected private IP to be rejected")
}
}
+74
View File
@@ -0,0 +1,74 @@
package mail
import (
"crypto/tls"
"encoding/base64"
"fmt"
"net/smtp"
"strings"
)
// SMTPConfig holds all the configuration parameters required to send an email.
type SMTPConfig struct {
Server string
Port int
Account string
Token string
SystemName string
}
// SendEmail sends an HTML email to the receiver using the provided SMTP configuration.
func SendEmail(config SMTPConfig, subject string, receiver string, content string) error {
encodedSubject := fmt.Sprintf("=?UTF-8?B?%s?=", base64.StdEncoding.EncodeToString([]byte(subject)))
mail := []byte(fmt.Sprintf("To: %s\r\n"+
"From: %s<%s>\r\n"+
"Subject: %s\r\n"+
"Content-Type: text/html; charset=UTF-8\r\n\r\n%s\r\n",
receiver, config.SystemName, config.Account, encodedSubject, content))
auth := smtp.PlainAuth("", config.Account, config.Token, config.Server)
addr := fmt.Sprintf("%s:%d", config.Server, config.Port)
to := strings.Split(receiver, ";")
var err error
if config.Port == 465 {
tlsConfig := &tls.Config{
InsecureSkipVerify: true,
ServerName: config.Server,
}
conn, err := tls.Dial("tcp", fmt.Sprintf("%s:%d", config.Server, config.Port), tlsConfig)
if err != nil {
return err
}
client, err := smtp.NewClient(conn, config.Server)
if err != nil {
return err
}
defer client.Close()
if err = client.Auth(auth); err != nil {
return err
}
if err = client.Mail(config.Account); err != nil {
return err
}
receiverEmails := strings.Split(receiver, ";")
for _, r := range receiverEmails {
if err = client.Rcpt(r); err != nil {
return err
}
}
w, err := client.Data()
if err != nil {
return err
}
_, err = w.Write(mail)
if err != nil {
return err
}
err = w.Close()
if err != nil {
return err
}
} else {
err = smtp.SendMail(addr, auth, config.Account, to, mail)
}
return err
}
+34
View File
@@ -0,0 +1,34 @@
package utils
import (
"log/slog"
"net"
"strings"
)
func GetIp() (ip string) {
ips, err := net.InterfaceAddrs()
if err != nil {
slog.Error("get interface addresses failed", "error", err)
return ip
}
for _, a := range ips {
if ipNet, ok := a.(*net.IPNet); ok && !ipNet.IP.IsLoopback() {
if ipNet.IP.To4() != nil {
ip = ipNet.IP.String()
if strings.HasPrefix(ip, "10") {
return
}
if strings.HasPrefix(ip, "172") {
return
}
if strings.HasPrefix(ip, "192.168") {
return
}
ip = ""
}
}
}
return
}
@@ -0,0 +1,67 @@
package ratelimit
import (
"sync"
"time"
)
type InMemoryRateLimiter struct {
store map[string]*[]int64
mutex sync.Mutex
expirationDuration time.Duration
}
func (l *InMemoryRateLimiter) Init(expirationDuration time.Duration) {
if l.store == nil {
l.mutex.Lock()
if l.store == nil {
l.store = make(map[string]*[]int64)
l.expirationDuration = expirationDuration
if expirationDuration > 0 {
go l.clearExpiredItems()
}
}
l.mutex.Unlock()
}
}
func (l *InMemoryRateLimiter) clearExpiredItems() {
for {
time.Sleep(l.expirationDuration)
l.mutex.Lock()
now := time.Now().Unix()
for key := range l.store {
queue := l.store[key]
size := len(*queue)
if size == 0 || now-(*queue)[size-1] > int64(l.expirationDuration.Seconds()) {
delete(l.store, key)
}
}
l.mutex.Unlock()
}
}
// Request parameter duration's unit is seconds
func (l *InMemoryRateLimiter) Request(key string, maxRequestNum int, duration int64) bool {
l.mutex.Lock()
defer l.mutex.Unlock()
// [old <-- new]
queue, ok := l.store[key]
now := time.Now().Unix()
if ok {
if len(*queue) < maxRequestNum {
*queue = append(*queue, now)
return true
}
if now-(*queue)[0] >= duration {
*queue = (*queue)[1:]
*queue = append(*queue, now)
return true
}
return false
}
s := make([]int64, 0, maxRequestNum)
l.store[key] = &s
*(l.store[key]) = append(*(l.store[key]), now)
return true
}
@@ -0,0 +1,994 @@
package openresty
import (
"crypto/sha256"
"crypto/x509"
"encoding/base64"
"encoding/hex"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"net/url"
"path"
"regexp"
"sort"
"strings"
)
func RenderJSON(sourceJSON string, certificateFiles []SupportFile) (*Result, error) {
var doc Document
if err := json.Unmarshal([]byte(strings.TrimSpace(sourceJSON)), &doc); err != nil {
return nil, fmt.Errorf("openresty source config json is invalid: %w", err)
}
return Render(doc, certificateFiles)
}
func Render(doc Document, certificateFiles []SupportFile) (*Result, error) {
mainConfig := RenderMainConfig(doc.OpenRestyConfig)
routeConfig, err := RenderRouteConfig(doc, certificateFiles)
if err != nil {
return nil, err
}
wafConfig, err := RenderWAFConfig(doc.WAF)
if err != nil {
return nil, err
}
files := append([]SupportFile(nil), certificateFiles...)
files = append(files, SupportFile{Path: "waf_config.json", Content: wafConfig})
files = DedupeSupportFiles(files)
return &Result{
MainConfig: mainConfig,
RouteConfig: routeConfig,
SupportFiles: files,
Checksum: ChecksumBundle(mainConfig, routeConfig, files),
}, nil
}
func RenderMainConfig(cfg ConfigSnapshot) string {
templateText := cfg.MainConfigTemplate
if strings.TrimSpace(templateText) == "" {
templateText = defaultMainConfigTemplate
}
return renderMainConfigTemplate(templateText, cfg)
}
func ValidateMainConfigTemplate(templateText string) error {
trimmed := strings.TrimSpace(templateText)
if trimmed == "" {
return errors.New("OpenRestyMainConfigTemplate 不能为空")
}
for _, placeholder := range requiredMainConfigTemplatePlaceholders {
if !strings.Contains(trimmed, placeholder) {
return fmt.Errorf("OpenRestyMainConfigTemplate 必须保留占位符 %s", placeholder)
}
}
return nil
}
func RenderRouteConfig(doc Document, certificateFiles []SupportFile) (string, error) {
var builder strings.Builder
builder.WriteString("# This file is generated by OpenFlare. Do not edit manually.\n")
certificates := certificatesByID(certificateFiles)
for _, route := range doc.Routes {
domains := normalizedRouteDomains(route)
if len(domains) == 0 {
return "", fmt.Errorf("route %s domains are invalid", route.Domain)
}
serverNames := renderServerNames(domains)
displayName := strings.TrimSpace(route.SiteName)
if displayName == "" {
displayName = domains[0]
}
cacheConfig := routeCacheConfig{Enabled: route.CacheEnabled, Policy: route.CachePolicy, Rules: route.CacheRules}
limitConfig := routeLimitConfig{LimitConnPerServer: route.LimitConnPerServer, LimitConnPerIP: route.LimitConnPerIP, LimitRate: route.LimitRate}
powEnabled, _ := getPoWConfigForRoute(route.ID, doc.WAF)
if normalizeRouteUpstreamType(route.UpstreamType) == "pages" {
if route.PagesDeployment == nil {
return "", fmt.Errorf("route %s pages deployment is missing", route.Domain)
}
if !route.EnableHTTPS {
builder.WriteString(renderHTTPPagesServer(serverNames, displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword))
continue
}
certIDs := normalizeCertIDs(route.CertID, route.CertIDs)
domainCertIDs := normalizeDomainCertIDs(domains, certIDs, route.DomainCertIDs)
if len(certIDs) == 0 {
return "", fmt.Errorf("路由 %s 未配置证书", route.Domain)
}
httpOnlyDomains := make([]string, 0, len(domains))
domainsByCertID := make(map[uint][]string, len(certIDs))
for index, domain := range domains {
if index >= len(domainCertIDs) || domainCertIDs[index] == 0 {
httpOnlyDomains = append(httpOnlyDomains, domain)
continue
}
domainsByCertID[domainCertIDs[index]] = append(domainsByCertID[domainCertIDs[index]], domain)
}
for _, certID := range certIDs {
assignedDomains := domainsByCertID[certID]
if len(assignedDomains) == 0 {
continue
}
certPEM, ok := certificates[certID]
if !ok {
return "", fmt.Errorf("route %s certificate %d does not exist", route.Domain, certID)
}
if err := validateCertificateCoverage(certPEM, assignedDomains); err != nil {
return "", fmt.Errorf("site %s certificate validation failed: %w", displayName, err)
}
}
if route.RedirectHTTP {
if len(httpOnlyDomains) > 0 {
builder.WriteString(renderHTTPPagesServer(renderServerNames(httpOnlyDomains), displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword))
}
for _, certID := range certIDs {
if assignedDomains := domainsByCertID[certID]; len(assignedDomains) > 0 {
builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains)))
}
}
} else {
builder.WriteString(renderHTTPPagesServer(serverNames, displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword))
}
for _, certID := range certIDs {
if assignedDomains := domainsByCertID[certID]; len(assignedDomains) > 0 {
builder.WriteString(renderHTTPSPagesServer(renderServerNames(assignedDomains), displayName, certID, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
}
}
continue
}
upstreams := route.Upstreams
if len(upstreams) == 0 && strings.TrimSpace(route.OriginURL) != "" {
upstreams = []string{route.OriginURL}
}
upstreamConfig := buildRouteUpstreamConfig(route, upstreams)
if upstreamConfig.UsesNamedUpstream {
builder.WriteString(renderNamedUpstreamBlock(upstreamConfig))
}
if !route.EnableHTTPS {
builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
continue
}
certIDs := normalizeCertIDs(route.CertID, route.CertIDs)
domainCertIDs := normalizeDomainCertIDs(domains, certIDs, route.DomainCertIDs)
if len(certIDs) == 0 {
return "", fmt.Errorf("路由 %s 未配置证书", route.Domain)
}
httpOnlyDomains := make([]string, 0, len(domains))
domainsByCertID := make(map[uint][]string, len(certIDs))
for index, domain := range domains {
if index >= len(domainCertIDs) || domainCertIDs[index] == 0 {
httpOnlyDomains = append(httpOnlyDomains, domain)
continue
}
domainsByCertID[domainCertIDs[index]] = append(domainsByCertID[domainCertIDs[index]], domain)
}
for _, certID := range certIDs {
assignedDomains := domainsByCertID[certID]
if len(assignedDomains) == 0 {
continue
}
certPEM, ok := certificates[certID]
if !ok {
return "", fmt.Errorf("route %s certificate %d does not exist", route.Domain, certID)
}
if err := validateCertificateCoverage(certPEM, assignedDomains); err != nil {
return "", fmt.Errorf("site %s certificate validation failed: %w", displayName, err)
}
}
if route.RedirectHTTP {
if len(httpOnlyDomains) > 0 {
builder.WriteString(renderHTTPProxyServer(renderServerNames(httpOnlyDomains), displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
}
for _, certID := range certIDs {
if assignedDomains := domainsByCertID[certID]; len(assignedDomains) > 0 {
builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains)))
}
}
} else {
builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
}
for _, certID := range certIDs {
if assignedDomains := domainsByCertID[certID]; len(assignedDomains) > 0 {
builder.WriteString(renderHTTPSServer(renderServerNames(assignedDomains), displayName, route.OriginURL, route.OriginHost, certID, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
}
}
}
return builder.String(), nil
}
func RenderPoWConfig(doc Document) (string, error) {
type domainEntry struct {
Domains []string `json:"domains"`
Enabled bool `json:"enabled"`
Config *PoWConfig `json:"config"`
}
entries := make([]domainEntry, 0)
for _, route := range doc.Routes {
powEnabled, powConfig := getPoWConfigForRoute(route.ID, doc.WAF)
if !powEnabled {
continue
}
entries = append(entries, domainEntry{Domains: normalizedRouteDomains(route), Enabled: true, Config: powConfig})
}
if len(entries) == 0 {
return "{}", nil
}
data, err := json.Marshal(entries)
return string(data), err
}
func RenderWAFConfig(snapshot WAFDocument) (string, error) {
type wafRuntimeRuleGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
IsGlobal bool `json:"is_global"`
BlockStatusCode int `json:"block_status_code"`
BlockResponseBody string `json:"block_response_body"`
IPWhitelist []string `json:"ip_whitelist"`
IPBlacklist []string `json:"ip_blacklist"`
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
CountryWhitelist []string `json:"country_whitelist"`
CountryBlacklist []string `json:"country_blacklist"`
RegionWhitelist []string `json:"region_whitelist"`
RegionBlacklist []string `json:"region_blacklist"`
PoWEnabled bool `json:"pow_enabled"`
PoWConfig *PoWConfig `json:"pow_config,omitempty"`
}
type wafRuntimeConfig struct {
DefaultBlockStatusCode int `json:"default_block_status_code"`
RuleGroups []wafRuntimeRuleGroup `json:"rule_groups"`
SiteRuleGroups map[string][]uint `json:"site_rule_groups"`
}
groups := make([]wafRuntimeRuleGroup, 0, len(snapshot.RuleGroups))
globalGroupIDs := make([]uint, 0)
enabledGroupIDs := make(map[uint]struct{}, len(snapshot.RuleGroups))
for _, group := range snapshot.RuleGroups {
if !group.Enabled {
continue
}
statusCode := group.BlockStatusCode
if statusCode == 0 {
statusCode = defaultWAFBlockStatus
}
if group.IsGlobal {
globalGroupIDs = append(globalGroupIDs, group.ID)
}
enabledGroupIDs[group.ID] = struct{}{}
powConfig := group.PoWConfig
if !group.PoWEnabled {
powConfig = nil
}
groups = append(groups, wafRuntimeRuleGroup{
ID: group.ID,
Name: group.Name,
IsGlobal: group.IsGlobal,
BlockStatusCode: statusCode,
BlockResponseBody: group.BlockResponseBody,
IPWhitelist: sortedUniqueStrings(group.IPWhitelist),
IPBlacklist: sortedUniqueStrings(group.IPBlacklist),
IPWhitelistGroups: sortedUniqueUintIDs(group.IPWhitelistGroups),
IPBlacklistGroups: sortedUniqueUintIDs(group.IPBlacklistGroups),
CountryWhitelist: group.CountryWhitelist,
CountryBlacklist: group.CountryBlacklist,
RegionWhitelist: group.RegionWhitelist,
RegionBlacklist: group.RegionBlacklist,
PoWEnabled: group.PoWEnabled,
PoWConfig: powConfig,
})
}
sort.Slice(groups, func(i, j int) bool {
if groups[i].IsGlobal != groups[j].IsGlobal {
return groups[i].IsGlobal
}
return groups[i].ID < groups[j].ID
})
sort.Slice(globalGroupIDs, func(i, j int) bool { return globalGroupIDs[i] < globalGroupIDs[j] })
siteRuleGroups := make(map[string][]uint, len(snapshot.Bindings))
for _, binding := range snapshot.Bindings {
ids := append([]uint{}, globalGroupIDs...)
for _, id := range binding.RuleGroupIDs {
if _, ok := enabledGroupIDs[id]; ok {
ids = append(ids, id)
}
}
siteRuleGroups[binding.SiteName] = uniqueUintIDs(ids)
}
data, err := json.Marshal(wafRuntimeConfig{DefaultBlockStatusCode: defaultWAFBlockStatus, RuleGroups: groups, SiteRuleGroups: siteRuleGroups})
return string(data), err
}
func sortedUniqueStrings(values []string) []string {
items := append([]string{}, values...)
items = uniqueStrings(items)
sort.Strings(items)
return items
}
func sortedUniqueUintIDs(values []uint) []uint {
items := uniqueUintIDs(values)
sort.Slice(items, func(i, j int) bool { return items[i] < items[j] })
return items
}
func ChecksumBundle(mainConfig string, routeConfig string, supportFiles []SupportFile) string {
var builder strings.Builder
builder.WriteString(mainConfig)
builder.WriteString("\n--route-config--\n")
builder.WriteString(routeConfig)
builder.WriteString("\n--support-files--\n")
files := DedupeSupportFiles(supportFiles)
sort.Slice(files, func(i int, j int) bool { return files[i].Path < files[j].Path })
for _, file := range files {
if file.Path == SourceConfigFileName {
continue
}
builder.WriteString(file.Path)
builder.WriteString("\n")
builder.WriteString(file.Content)
builder.WriteString("\n")
}
sum := sha256.Sum256([]byte(builder.String()))
return hex.EncodeToString(sum[:])
}
func DedupeSupportFiles(files []SupportFile) []SupportFile {
if len(files) == 0 {
return nil
}
unique := make(map[string]SupportFile, len(files))
for _, file := range files {
unique[file.Path] = file
}
result := make([]SupportFile, 0, len(unique))
for _, file := range unique {
result = append(result, file)
}
return result
}
func renderMainConfigTemplate(templateText string, cfg ConfigSnapshot) string {
replacer := strings.NewReplacer(
"{{OpenRestyWorkerProcesses}}", cfg.WorkerProcesses,
"{{OpenRestyWorkerConnections}}", fmt.Sprintf("%d", cfg.WorkerConnections),
"{{OpenRestyWorkerRlimitNofile}}", fmt.Sprintf("%d", cfg.WorkerRlimitNofile),
"{{OpenRestyConnectionUpgradeMap}}", renderConnectionUpgradeMap(),
"{{OpenRestyDefaultServerBlock}}", renderDefaultServerBlock(cfg.DefaultServerReturnStatus, cfg.HTTP3Enabled),
"{{OpenRestyAccessLogPath}}", AccessLogPlaceholder,
"{{OpenRestyErrorLogPath}}", ErrorLogPlaceholder,
"{{OpenRestyEventsUseDirective}}", renderTemplateDirective(cfg.EventsUse != "", fmt.Sprintf("use %s;", cfg.EventsUse)),
"{{OpenRestyEventsMultiAcceptDirective}}", renderTemplateDirective(cfg.EventsMultiAcceptEnabled, "multi_accept on;"),
"{{OpenRestyKeepaliveTimeout}}", fmt.Sprintf("%d", cfg.KeepaliveTimeout),
"{{OpenRestyKeepaliveRequests}}", fmt.Sprintf("%d", cfg.KeepaliveRequests),
"{{OpenRestyClientHeaderTimeout}}", fmt.Sprintf("%d", cfg.ClientHeaderTimeout),
"{{OpenRestyClientBodyTimeout}}", fmt.Sprintf("%d", cfg.ClientBodyTimeout),
"{{OpenRestyClientMaxBodySize}}", cfg.ClientMaxBodySize,
"{{OpenRestyLargeClientHeaderBuffers}}", cfg.LargeClientHeaderBuffers,
"{{OpenRestySendTimeout}}", fmt.Sprintf("%d", cfg.SendTimeout),
"{{OpenRestyProxyConnectTimeout}}", fmt.Sprintf("%d", cfg.ProxyConnectTimeout),
"{{OpenRestyProxySendTimeout}}", fmt.Sprintf("%d", cfg.ProxySendTimeout),
"{{OpenRestyProxyReadTimeout}}", fmt.Sprintf("%d", cfg.ProxyReadTimeout),
"{{OpenRestyProxyRequestBuffering}}", onOff(cfg.ProxyRequestBuffering),
"{{OpenRestyProxyBuffering}}", onOff(cfg.ProxyBufferingEnabled),
"{{OpenRestyProxyBuffers}}", cfg.ProxyBuffers,
"{{OpenRestyProxyBufferSize}}", cfg.ProxyBufferSize,
"{{OpenRestyProxyBusyBuffersSize}}", cfg.ProxyBusyBuffersSize,
"{{OpenRestyGzip}}", onOff(cfg.GzipEnabled),
"{{OpenRestyGzipMinLength}}", fmt.Sprintf("%d", cfg.GzipMinLength),
"{{OpenRestyGzipCompLevel}}", fmt.Sprintf("%d", cfg.GzipCompLevel),
"{{OpenRestyResolverDirective}}", renderTemplateDirective(cfg.Resolvers != "", fmt.Sprintf("resolver %s;", cfg.Resolvers)),
"{{OpenRestyCacheBlock}}", renderOpenRestyCacheTemplateBlock(cfg),
"{{OpenRestyRouteConfigInclude}}", RouteConfigPlaceholder,
)
return replacer.Replace(templateText)
}
func renderTemplateDirective(enabled bool, statement string) string {
if !enabled {
return ""
}
return fmt.Sprintf(" %s\n", statement)
}
func renderOpenRestyCacheTemplateBlock(cfg ConfigSnapshot) string {
lines := []string{renderOpenRestyLimitZoneBlock()}
if !cfg.CacheEnabled {
lines = append(lines, renderOpenRestyObservabilityTemplateBlock())
return strings.Join(lines, "")
}
lines = append(lines, strings.Join([]string{
fmt.Sprintf(" proxy_cache_path %s levels=%s keys_zone=openflare_cache:10m inactive=%s max_size=%s;", cfg.CachePath, cfg.CacheLevels, cfg.CacheInactive, cfg.CacheMaxSize),
fmt.Sprintf(" proxy_cache_key \"%s\";", cfg.CacheKeyTemplate),
fmt.Sprintf(" proxy_cache_lock %s;", onOff(cfg.CacheLockEnabled)),
fmt.Sprintf(" proxy_cache_lock_timeout %s;", cfg.CacheLockTimeout),
fmt.Sprintf(" proxy_cache_use_stale %s;", cfg.CacheUseStale),
"",
}, "\n"))
lines = append(lines, renderOpenRestyObservabilityTemplateBlock())
return strings.Join(lines, "")
}
func renderOpenRestyLimitZoneBlock() string {
return " limit_conn_zone $server_name zone=openflare_conn_per_server:10m;\n limit_conn_zone $binary_remote_addr zone=openflare_conn_per_ip:10m;\n"
}
func renderOpenRestyObservabilityTemplateBlock() string {
return fmt.Sprintf(" lua_shared_dict openflare_observability 10m;\n lua_shared_dict openflare_pow_challenges 10m;\n lua_shared_dict openflare_pow_sessions 10m;\n lua_shared_dict openflare_pow_config 1m;\n lua_shared_dict openflare_waf_config 1m;\n init_worker_by_lua_file %s/observability/init.lua;\n log_by_lua_file %s/observability/log.lua;\n\n server {\n listen %s;\n server_name openflare-observability;\n access_log off;\n\n location = /openflare/stub_status {\n stub_status;\n }\n\n location = /openflare/observability {\n default_type application/json;\n content_by_lua_file %s/observability/read.lua;\n }\n }\n\n", LuaDirPlaceholder, LuaDirPlaceholder, ObservabilityListenPlaceholder, LuaDirPlaceholder)
}
func renderHTTPProxyServer(serverNames string, siteName string, originURL string, originHost string, customHeaders []CustomHeader, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg ConfigSnapshot) string {
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n%s%s location / {\n%s%s%s%s%s }\n%s}\n\n", serverNames, renderAccessBlock(siteName, powEnabled), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig, cfg), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig), renderPowStaticLocationBlock(powEnabled))
}
func renderPagesAPIProxyLocationBlock(deployment *PagesDeployment) string {
if deployment == nil || !deployment.APIProxyEnabled {
return ""
}
path := strings.TrimSpace(deployment.APIProxyPath)
pass := strings.TrimSpace(deployment.APIProxyPass)
rewrite := strings.TrimSpace(deployment.APIProxyRewrite)
if path == "" || pass == "" {
return ""
}
if !strings.HasPrefix(path, "/") {
path = "/" + path
}
cleanPath := strings.TrimSuffix(path, "/")
var builder strings.Builder
builder.WriteString(fmt.Sprintf("\n location %s {\n", cleanPath))
if rewrite != "" {
if !strings.HasPrefix(rewrite, "/") {
rewrite = "/" + rewrite
}
cleanRewrite := strings.TrimSuffix(rewrite, "/")
if cleanRewrite == "" {
builder.WriteString(fmt.Sprintf(" rewrite ^%s/(.*)$ /$1 break;\n", regexp.QuoteMeta(cleanPath)))
builder.WriteString(fmt.Sprintf(" rewrite ^%s$ / break;\n", regexp.QuoteMeta(cleanPath)))
} else {
builder.WriteString(fmt.Sprintf(" rewrite ^%s/(.*)$ %s/$1 break;\n", regexp.QuoteMeta(cleanPath), cleanRewrite))
builder.WriteString(fmt.Sprintf(" rewrite ^%s$ %s break;\n", regexp.QuoteMeta(cleanPath), cleanRewrite))
}
}
builder.WriteString(fmt.Sprintf(" proxy_pass %s;\n", pass))
builder.WriteString(" proxy_http_version 1.1;\n")
builder.WriteString(" proxy_set_header Host $http_host;\n")
builder.WriteString(" proxy_set_header X-Real-IP $remote_addr;\n")
builder.WriteString(" proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;\n")
builder.WriteString(" proxy_set_header X-Forwarded-Proto $scheme;\n")
builder.WriteString(" proxy_set_header Upgrade $http_upgrade;\n")
builder.WriteString(" proxy_set_header Connection $connection_upgrade;\n")
builder.WriteString(" }\n")
return builder.String()
}
func renderHTTPPagesServer(serverNames string, siteName string, deployment *PagesDeployment, limitConfig routeLimitConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string) string {
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n%s%s root %s;\n index %s;%s\n\n location / {\n%s%s }\n%s}\n\n", serverNames, renderAccessBlock(siteName, powEnabled), renderPowLocationBlocks(powEnabled), quoteNginxStringLiteral(pagesDeploymentRoot(deployment)), quoteNginxStringLiteral(pagesEntryFile(deployment)), renderPagesAPIProxyLocationBlock(deployment), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderPagesLocationBlock(deployment, limitConfig), renderPowStaticLocationBlock(powEnabled))
}
func renderHTTPRedirectServer(serverNames string) string {
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n\n return 301 https://$host$request_uri;\n}\n\n", serverNames)
}
func renderHTTPSServer(serverNames string, siteName string, originURL string, originHost string, certificateID uint, customHeaders []CustomHeader, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg ConfigSnapshot) string {
certPath := fmt.Sprintf("%s/%d.crt", CertDirPlaceholder, certificateID)
keyPath := fmt.Sprintf("%s/%d.key", CertDirPlaceholder, certificateID)
var h3Listen string
var h3Header string
if cfg.HTTP3Enabled {
h3Listen = " listen 443 quic;\n"
h3Header = " add_header Alt-Svc 'h3=\":443\"; ma=86400';\n"
}
return fmt.Sprintf("server {\n listen 443 ssl;\n%s http2 on;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n%s%s%s location / {\n%s%s%s%s%s }\n%s}\n\n", h3Listen, serverNames, certPath, keyPath, h3Header, renderAccessBlock(siteName, powEnabled), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig, cfg), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig), renderPowStaticLocationBlock(powEnabled))
}
func renderHTTPSPagesServer(serverNames string, siteName string, certificateID uint, deployment *PagesDeployment, limitConfig routeLimitConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg ConfigSnapshot) string {
certPath := fmt.Sprintf("%s/%d.crt", CertDirPlaceholder, certificateID)
keyPath := fmt.Sprintf("%s/%d.key", CertDirPlaceholder, certificateID)
var h3Listen string
var h3Header string
if cfg.HTTP3Enabled {
h3Listen = " listen 443 quic;\n"
h3Header = " add_header Alt-Svc 'h3=\":443\"; ma=86400';\n"
}
return fmt.Sprintf("server {\n listen 443 ssl;\n%s http2 on;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n%s%s%s root %s;\n index %s;%s\n\n location / {\n%s%s }\n%s}\n\n", h3Listen, serverNames, certPath, keyPath, h3Header, renderAccessBlock(siteName, powEnabled), renderPowLocationBlocks(powEnabled), quoteNginxStringLiteral(pagesDeploymentRoot(deployment)), quoteNginxStringLiteral(pagesEntryFile(deployment)), renderPagesAPIProxyLocationBlock(deployment), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderPagesLocationBlock(deployment, limitConfig), renderPowStaticLocationBlock(powEnabled))
}
func renderPagesLocationBlock(deployment *PagesDeployment, limitConfig routeLimitConfig) string {
var builder strings.Builder
builder.WriteString(renderRouteLimitBlock(limitConfig))
if deployment != nil && deployment.SPAFallbackEnabled {
builder.WriteString(fmt.Sprintf(" try_files $uri $uri/ %s;\n", pagesFallbackPath(deployment)))
} else {
builder.WriteString(" try_files $uri $uri/ =404;\n")
}
return builder.String()
}
func pagesDeploymentRoot(deployment *PagesDeployment) string {
if deployment == nil || strings.TrimSpace(deployment.LocalRoot) == "" {
return PagesDirPlaceholder
}
return filepathToNginxPath(deployment.LocalRoot)
}
func pagesEntryFile(deployment *PagesDeployment) string {
if deployment == nil || strings.TrimSpace(deployment.EntryFile) == "" {
return "index.html"
}
return strings.TrimPrefix(filepathToNginxPath(deployment.EntryFile), "/")
}
func pagesFallbackPath(deployment *PagesDeployment) string {
if deployment == nil || strings.TrimSpace(deployment.SPAFallbackPath) == "" {
return "/index.html"
}
value := filepathToNginxPath(strings.TrimSpace(deployment.SPAFallbackPath))
if !strings.HasPrefix(value, "/") {
value = "/" + value
}
if value == "/" || strings.HasSuffix(value, "/") || strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") || strings.ContainsAny(value, " \t\r\n") {
return "/index.html"
}
for _, segment := range strings.Split(value, "/") {
if segment == "." || segment == ".." {
return "/index.html"
}
}
cleaned := path.Clean(value)
if cleaned == "/" || strings.HasSuffix(cleaned, "/") {
return "/index.html"
}
return cleaned
}
func renderProxyHeaderBlock(originURL string, originHost string, customHeaders []CustomHeader, upstreamConfig routeUpstreamConfig, cfg ConfigSnapshot) string {
var builder strings.Builder
if strings.TrimSpace(originHost) != "" {
builder.WriteString(fmt.Sprintf(" proxy_set_header Host %s;\n", quoteNginxStringLiteral(originHost)))
} else {
builder.WriteString(" proxy_set_header Host $host;\n")
}
if upstreamServerName := resolveUpstreamServerName(originURL, originHost); upstreamServerName != "" {
builder.WriteString(" proxy_ssl_server_name on;\n")
builder.WriteString(fmt.Sprintf(" proxy_ssl_name %s;\n", quoteNginxStringLiteral(upstreamServerName)))
}
builder.WriteString(" proxy_set_header X-Real-IP $remote_addr;\n")
builder.WriteString(" proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;\n")
builder.WriteString(" proxy_set_header X-Forwarded-Proto $scheme;\n")
if cfg.WebsocketEnabled {
builder.WriteString(" proxy_http_version 1.1;\n")
builder.WriteString(" proxy_set_header Connection $connection_upgrade;\n")
builder.WriteString(" proxy_set_header Upgrade $http_upgrade;\n")
} else if upstreamConfig.UsesNamedUpstream {
builder.WriteString(" proxy_http_version 1.1;\n")
builder.WriteString(" proxy_set_header Connection \"\";\n")
}
for _, header := range customHeaders {
builder.WriteString(fmt.Sprintf(" proxy_set_header %s %s;\n", header.Key, quoteNginxStringLiteral(header.Value)))
}
return builder.String()
}
func renderAccessBlock(siteName string, powEnabled bool) string {
escapedSiteName := escapeNginxString(siteName)
if !powEnabled {
return fmt.Sprintf(" set $openflare_waf_site \"%s\";\n access_by_lua_file %s/waf/check.lua;\n", escapedSiteName, LuaDirPlaceholder)
}
return fmt.Sprintf(` set $openflare_waf_site "%s";
access_by_lua_block {
if not string.find(package.path, "%s/?.lua", 1, true) then
package.path = "%s/?.lua;%s/?/init.lua;" .. package.path
end
require("waf.runtime").check()
if ngx.ctx.openflare_waf_blocked then
return
end
require("pow.runtime").check()
}
`, escapedSiteName, LuaDirPlaceholder, LuaDirPlaceholder, LuaDirPlaceholder)
}
func renderBasicAuthBlock(enabled bool, username, password string) string {
if !enabled || username == "" || password == "" {
return ""
}
encoded := base64.StdEncoding.EncodeToString([]byte(username + ":" + password))
return fmt.Sprintf(" rewrite_by_lua_block {\n local auth = ngx.var.http_authorization\n if auth ~= \"Basic %s\" then\n ngx.header[\"WWW-Authenticate\"] = 'Basic realm=\"Restricted\"'\n return ngx.exit(401)\n end\n }\n", encoded)
}
func renderPowLocationBlocks(powEnabled bool) string {
if !powEnabled {
return ""
}
return fmt.Sprintf("\n location = %spass-challenge {\n content_by_lua_file %s/pow/verify.lua;\n }\n\n location = %smake-challenge {\n content_by_lua_file %s/pow/challenge.lua;\n }\n\n", anubisAPIPrefix, LuaDirPlaceholder, anubisAPIPrefix, LuaDirPlaceholder)
}
func renderPowStaticLocationBlock(powEnabled bool) string {
if !powEnabled {
return ""
}
return fmt.Sprintf(" location %s {\n alias %s/;\n types {\n text/css css;\n application/javascript js mjs;\n application/json json;\n image/webp webp;\n font/woff2 woff2;\n }\n }\n\n", anubisStaticPrefix, PowStaticDirPlaceholder)
}
func renderRouteCacheBlock(cacheConfig routeCacheConfig, cfg ConfigSnapshot) string {
if !cfg.CacheEnabled || !cacheConfig.Enabled {
return ""
}
var builder strings.Builder
builder.WriteString(" set $openflare_skip_cache 0;\n")
builder.WriteString(" if ($request_method != GET) {\n set $openflare_skip_cache 1;\n }\n")
builder.WriteString(" if ($http_authorization != \"\") {\n set $openflare_skip_cache 1;\n }\n")
builder.WriteString(" if ($http_cookie ~* \"(session|sess|token|auth|jwt|logged_in|remember|laravel_session|connect\\\\.sid|_session)\") {\n set $openflare_skip_cache 1;\n }\n")
builder.WriteString(" if ($http_cache_control ~* \"(no-cache|no-store|private)\") {\n set $openflare_skip_cache 1;\n }\n")
if condition := renderRouteCachePolicyCondition(cacheConfig); condition != "" {
builder.WriteString(condition)
}
builder.WriteString(" proxy_cache openflare_cache;\n")
builder.WriteString(" proxy_cache_methods GET;\n")
builder.WriteString(" proxy_cache_bypass $openflare_skip_cache;\n")
builder.WriteString(" proxy_no_cache $openflare_skip_cache;\n")
return builder.String()
}
func renderRouteLimitBlock(limitConfig routeLimitConfig) string {
var builder strings.Builder
if limitConfig.LimitConnPerServer > 0 {
builder.WriteString(fmt.Sprintf(" limit_conn openflare_conn_per_server %d;\n", limitConfig.LimitConnPerServer))
}
if limitConfig.LimitConnPerIP > 0 {
builder.WriteString(fmt.Sprintf(" limit_conn openflare_conn_per_ip %d;\n", limitConfig.LimitConnPerIP))
}
if strings.TrimSpace(limitConfig.LimitRate) != "" {
builder.WriteString(fmt.Sprintf(" limit_rate %s;\n", limitConfig.LimitRate))
}
return builder.String()
}
func renderRouteCachePolicyCondition(cacheConfig routeCacheConfig) string {
switch cacheConfig.Policy {
case cachePolicySuffix:
return fmt.Sprintf(" if ($uri !~* %s) {\n set $openflare_skip_cache 1;\n }\n", quoteNginxStringLiteral(buildSuffixMatchPattern(cacheConfig.Rules)))
case cachePolicyPathPrefix:
return fmt.Sprintf(" if ($uri !~ %s) {\n set $openflare_skip_cache 1;\n }\n", quoteNginxStringLiteral(buildPathPrefixMatchPattern(cacheConfig.Rules)))
case cachePolicyPathExact:
return fmt.Sprintf(" if ($uri !~ %s) {\n set $openflare_skip_cache 1;\n }\n", quoteNginxStringLiteral(buildPathExactMatchPattern(cacheConfig.Rules)))
default:
return ""
}
}
func renderProxyPassBlock(originURL string, upstreamConfig routeUpstreamConfig) string {
parsed, err := url.Parse(originURL)
if err != nil || parsed.Host == "" || parsed.Scheme == "" {
return fmt.Sprintf(" proxy_pass %s;\n", originURL)
}
if upstreamConfig.UsesNamedUpstream {
return fmt.Sprintf(" proxy_pass %s://%s%s;\n", upstreamConfig.Scheme, upstreamConfig.Name, upstreamConfig.ProxyPassURI)
}
return fmt.Sprintf(" proxy_pass %s;\n", originURL)
}
func buildRouteUpstreamConfig(route Route, upstreams []string) routeUpstreamConfig {
if len(upstreams) == 0 {
return routeUpstreamConfig{}
}
if len(upstreams) == 1 {
parsed, err := url.Parse(strings.TrimSpace(upstreams[0]))
if err != nil || parsed.Host == "" || parsed.Scheme == "" {
return routeUpstreamConfig{}
}
return routeUpstreamConfig{Name: buildRouteUpstreamName(route), Scheme: parsed.Scheme, ProxyPassURI: buildUpstreamProxyPassURI(parsed), Servers: []string{parsed.Host}, UsesNamedUpstream: true}
}
servers := make([]string, 0, len(upstreams))
var scheme string
for _, upstream := range upstreams {
parsed, err := url.Parse(strings.TrimSpace(upstream))
if err != nil || parsed.Host == "" || parsed.Scheme == "" || (strings.TrimSpace(parsed.EscapedPath()) != "" && strings.TrimSpace(parsed.EscapedPath()) != "/") || parsed.RawQuery != "" {
return routeUpstreamConfig{}
}
if scheme == "" {
scheme = parsed.Scheme
} else if scheme != parsed.Scheme {
return routeUpstreamConfig{}
}
servers = append(servers, parsed.Host)
}
return routeUpstreamConfig{Name: buildRouteUpstreamName(route), Scheme: scheme, Servers: servers, UsesNamedUpstream: true}
}
func normalizeRouteUpstreamType(raw string) string {
switch strings.ToLower(strings.TrimSpace(raw)) {
case "pages":
return "pages"
default:
return "direct"
}
}
func renderNamedUpstreamBlock(upstreamConfig routeUpstreamConfig) string {
var builder strings.Builder
builder.WriteString(fmt.Sprintf("upstream %s {\n", upstreamConfig.Name))
for _, server := range upstreamConfig.Servers {
builder.WriteString(fmt.Sprintf(" server %s max_fails=3 fail_timeout=10s;\n", server))
}
builder.WriteString(" keepalive 128;\n}\n\n")
return builder.String()
}
func buildRouteUpstreamName(route Route) string {
identity := strings.TrimSpace(route.SiteName)
if identity == "" {
identity = route.Domain
}
sanitized := strings.Map(func(r rune) rune {
switch {
case r >= 'a' && r <= 'z':
return r
case r >= 'A' && r <= 'Z':
return r + ('a' - 'A')
case r >= '0' && r <= '9':
return r
default:
return '_'
}
}, identity)
sanitized = strings.Trim(sanitized, "_")
if sanitized == "" {
sanitized = "backend"
}
return fmt.Sprintf("backend_%s_%d", sanitized, route.ID)
}
func buildUpstreamProxyPassURI(parsed *url.URL) string {
path := parsed.EscapedPath()
if path == "/" {
path = ""
}
if parsed.RawQuery == "" {
return path
}
return fmt.Sprintf("%s?%s", path, parsed.RawQuery)
}
func renderConnectionUpgradeMap() string {
return " map $http_upgrade $connection_upgrade {\n default upgrade;\n '' \"\";\n }\n\n"
}
func renderDefaultServerBlock(statusCode int, http3Enabled bool) string {
if statusCode <= 0 {
statusCode = 421
}
var h3Default string
if http3Enabled {
h3Default = "\n listen 443 quic reuseport default_server;"
}
return strings.Join([]string{
" server {",
" listen 80 default_server;",
" server_name _;",
"",
fmt.Sprintf(" return %d;", statusCode),
" }",
"",
" server {",
fmt.Sprintf(" listen 443 ssl default_server;%s", h3Default),
" server_name _;",
"",
" ssl_reject_handshake on;",
" }",
"",
}, "\n")
}
func normalizedRouteDomains(route Route) []string {
if len(route.Domains) > 0 {
return route.Domains
}
if strings.TrimSpace(route.Domain) == "" {
return nil
}
return []string{strings.TrimSpace(route.Domain)}
}
func normalizeCertIDs(primaryCertID *uint, certIDs []uint) []uint {
candidates := make([]uint, 0, len(certIDs)+1)
if primaryCertID != nil && *primaryCertID != 0 {
candidates = append(candidates, *primaryCertID)
}
candidates = append(candidates, certIDs...)
seen := make(map[uint]struct{}, len(candidates))
normalized := make([]uint, 0, len(candidates))
for _, id := range candidates {
if id == 0 {
continue
}
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
normalized = append(normalized, id)
}
return normalized
}
func normalizeDomainCertIDs(domains []string, certIDs []uint, domainCertIDs []uint) []uint {
if len(domainCertIDs) > 0 {
normalized := make([]uint, len(domainCertIDs))
copy(normalized, domainCertIDs)
return normalized
}
if len(certIDs) == 1 {
normalized := make([]uint, len(domains))
for index := range normalized {
normalized[index] = certIDs[0]
}
return normalized
}
if len(certIDs) == len(domains) {
normalized := make([]uint, len(certIDs))
copy(normalized, certIDs)
return normalized
}
return []uint{}
}
func certificatesByID(files []SupportFile) map[uint]string {
result := make(map[uint]string)
for _, file := range files {
if !strings.HasSuffix(file.Path, ".crt") {
continue
}
idText := strings.TrimSuffix(file.Path, ".crt")
var id uint
if _, err := fmt.Sscanf(idText, "%d", &id); err == nil && id != 0 {
result[id] = file.Content
}
}
return result
}
func validateCertificateCoverage(certPEM string, domains []string) error {
block, _ := pem.Decode([]byte(certPEM))
if block == nil {
return errors.New("certificate PEM is invalid")
}
leaf, err := x509.ParseCertificate(block.Bytes)
if err != nil {
return err
}
for _, domain := range domains {
if err := leaf.VerifyHostname(domain); err != nil {
return fmt.Errorf("certificate does not cover domain %s", domain)
}
}
return nil
}
func getPoWConfigForRoute(routeID uint, snapshot WAFDocument) (bool, *PoWConfig) {
for _, binding := range snapshot.Bindings {
if binding.RouteID != routeID {
continue
}
for _, groupID := range binding.RuleGroupIDs {
for _, group := range snapshot.RuleGroups {
if group.ID == groupID && group.PoWEnabled {
return true, group.PoWConfig
}
}
}
break
}
for _, group := range snapshot.RuleGroups {
if group.IsGlobal && group.PoWEnabled {
return true, group.PoWConfig
}
}
return false, nil
}
func uniqueUintIDs(values []uint) []uint {
seen := make(map[uint]struct{}, len(values))
result := make([]uint, 0, len(values))
for _, value := range values {
if value == 0 {
continue
}
if _, ok := seen[value]; ok {
continue
}
seen[value] = struct{}{}
result = append(result, value)
}
return result
}
func uniqueStrings(values []string) []string {
seen := make(map[string]struct{}, len(values))
result := make([]string, 0, len(values))
for _, value := range values {
item := strings.TrimSpace(value)
if item == "" {
continue
}
if _, ok := seen[item]; ok {
continue
}
seen[item] = struct{}{}
result = append(result, item)
}
return result
}
func resolveUpstreamServerName(originURL string, originHost string) string {
parsed, err := url.Parse(originURL)
if err != nil || !strings.EqualFold(parsed.Scheme, "https") {
return ""
}
if strings.TrimSpace(originHost) != "" {
parsedHost, err := url.Parse("//" + originHost)
if err == nil && parsedHost.Hostname() != "" {
return parsedHost.Hostname()
}
return originHost
}
return parsed.Hostname()
}
func renderServerNames(domains []string) string { return strings.Join(domains, " ") }
func onOff(value bool) string {
if value {
return "on"
}
return "off"
}
func quoteNginxStringLiteral(value string) string {
escaped := strings.ReplaceAll(value, `\`, `\\`)
escaped = strings.ReplaceAll(escaped, `"`, `\"`)
return fmt.Sprintf(`"%s"`, escaped)
}
func filepathToNginxPath(value string) string {
return strings.ReplaceAll(strings.TrimSpace(value), `\`, `/`)
}
func escapeNginxString(value string) string {
escaped := strings.ReplaceAll(value, `\`, `\\`)
escaped = strings.ReplaceAll(escaped, `"`, `\"`)
return escaped
}
func buildSuffixMatchPattern(rules []string) string {
parts := make([]string, 0, len(rules))
for _, rule := range rules {
parts = append(parts, regexp.QuoteMeta(rule))
}
return fmt.Sprintf("\\.(?:%s)$", strings.Join(parts, "|"))
}
func buildPathPrefixMatchPattern(rules []string) string {
parts := make([]string, 0, len(rules))
for _, rule := range rules {
trimmed := strings.TrimRight(rule, "/")
if trimmed == "" {
trimmed = "/"
}
if trimmed == "/" {
parts = append(parts, "/")
continue
}
parts = append(parts, fmt.Sprintf("%s(?:/|$)", regexp.QuoteMeta(trimmed)))
}
return fmt.Sprintf("^(?:%s)", strings.Join(parts, "|"))
}
func buildPathExactMatchPattern(rules []string) string {
parts := make([]string, 0, len(rules))
for _, rule := range rules {
parts = append(parts, regexp.QuoteMeta(rule))
}
return fmt.Sprintf("^(?:%s)$", strings.Join(parts, "|"))
}
@@ -0,0 +1,100 @@
package openresty
import (
"strings"
"testing"
)
func TestRenderPagesAPIProxyLocationBlock(t *testing.T) {
tests := []struct {
name string
deployment *PagesDeployment
expected []string
unexpected []string
}{
{
name: "nil deployment",
deployment: nil,
expected: []string{""},
},
{
name: "disabled proxy",
deployment: &PagesDeployment{
APIProxyEnabled: false,
APIProxyPath: "/api",
APIProxyPass: "http://127.0.0.1:8080",
},
expected: []string{""},
},
{
name: "enabled proxy without rewrite",
deployment: &PagesDeployment{
APIProxyEnabled: true,
APIProxyPath: "/api",
APIProxyPass: "http://127.0.0.1:8080",
APIProxyRewrite: "",
},
expected: []string{
"location /api {",
"proxy_pass http://127.0.0.1:8080;",
"proxy_http_version 1.1;",
"proxy_set_header Host $http_host;",
},
unexpected: []string{
"rewrite",
},
},
{
name: "enabled proxy with rewrite to root",
deployment: &PagesDeployment{
APIProxyEnabled: true,
APIProxyPath: "/api",
APIProxyPass: "http://127.0.0.1:8080",
APIProxyRewrite: "/",
},
expected: []string{
"location /api {",
"rewrite ^/api/(.*)$ /$1 break;",
"rewrite ^/api$ / break;",
"proxy_pass http://127.0.0.1:8080;",
},
},
{
name: "enabled proxy with rewrite to subpath",
deployment: &PagesDeployment{
APIProxyEnabled: true,
APIProxyPath: "/api",
APIProxyPass: "http://127.0.0.1:8080",
APIProxyRewrite: "/v2",
},
expected: []string{
"location /api {",
"rewrite ^/api/(.*)$ /v2/$1 break;",
"rewrite ^/api$ /v2 break;",
"proxy_pass http://127.0.0.1:8080;",
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := renderPagesAPIProxyLocationBlock(tt.deployment)
if len(tt.expected) == 1 && tt.expected[0] == "" {
if got != "" {
t.Fatalf("expected empty output, got: %q", got)
}
return
}
for _, exp := range tt.expected {
if !strings.Contains(got, exp) {
t.Errorf("expected output to contain %q, but got:\n%s", exp, got)
}
}
for _, unexp := range tt.unexpected {
if strings.Contains(got, unexp) {
t.Errorf("expected output NOT to contain %q, but got:\n%s", unexp, got)
}
}
})
}
}
@@ -0,0 +1,282 @@
package openresty
const (
CertDirPlaceholder = "__OPENFLARE_CERT_DIR__"
RouteConfigPlaceholder = "__OPENFLARE_ROUTE_CONFIG__"
AccessLogPlaceholder = "__OPENFLARE_ACCESS_LOG__"
ErrorLogPlaceholder = "__OPENFLARE_ERROR_LOG__"
LuaDirPlaceholder = "__OPENFLARE_LUA_DIR__"
ObservabilityListenPlaceholder = "__OPENFLARE_OBSERVABILITY_LISTEN__"
ObservabilityPortPlaceholder = "__OPENFLARE_OBSERVABILITY_PORT__"
PowStaticDirPlaceholder = "__OPENFLARE_POW_STATIC_DIR__"
PagesDirPlaceholder = "__OPENFLARE_PAGES_DIR__"
SourceConfigFileName = "openresty_config.json"
)
const (
cachePolicySuffix = "suffix"
cachePolicyPathPrefix = "path_prefix"
cachePolicyPathExact = "path_exact"
defaultWAFBlockStatus = 418
anubisStaticPrefix = "/.within.website/x/cmd/anubis/static/"
anubisAPIPrefix = "/.within.website/x/cmd/anubis/api/"
)
const defaultMainConfigTemplate = `# 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}};
}
`
type SupportFile struct {
Path string `json:"path"`
Content string `json:"content"`
}
type CustomHeader struct {
Key string `json:"key"`
Value string `json:"value"`
}
type PoWListConfig struct {
IPs []string `json:"ips"`
IPCidrs []string `json:"ip_cidrs"`
Paths []string `json:"paths"`
PathRegexes []string `json:"path_regexes"`
UserAgents []string `json:"user_agents"`
}
type PoWConfig struct {
Difficulty int `json:"difficulty"`
Algorithm string `json:"algorithm"`
SessionTTL int `json:"session_ttl"`
ChallengeTTL int `json:"challenge_ttl"`
Whitelist PoWListConfig `json:"whitelist"`
Blacklist PoWListConfig `json:"blacklist"`
}
type Route struct {
ID uint `json:"id,omitempty"`
SiteName string `json:"site_name,omitempty"`
Domain string `json:"domain"`
Domains []string `json:"domains,omitempty"`
OriginURL string `json:"origin_url"`
OriginHost string `json:"origin_host,omitempty"`
Upstreams []string `json:"upstreams,omitempty"`
Enabled bool `json:"enabled"`
EnableHTTPS bool `json:"enable_https"`
CertID *uint `json:"cert_id,omitempty"`
CertIDs []uint `json:"cert_ids,omitempty"`
DomainCertIDs []uint `json:"domain_cert_ids,omitempty"`
RedirectHTTP bool `json:"redirect_http"`
LimitConnPerServer int `json:"limit_conn_per_server,omitempty"`
LimitConnPerIP int `json:"limit_conn_per_ip,omitempty"`
LimitRate string `json:"limit_rate,omitempty"`
CacheEnabled bool `json:"cache_enabled"`
CachePolicy string `json:"cache_policy,omitempty"`
CacheRules []string `json:"cache_rules,omitempty"`
CustomHeaders []CustomHeader `json:"custom_headers,omitempty"`
PoWEnabled bool `json:"pow_enabled,omitempty"`
PoWConfig *PoWConfig `json:"pow_config,omitempty"`
BasicAuthEnabled bool `json:"basic_auth_enabled,omitempty"`
BasicAuthUsername string `json:"basic_auth_username,omitempty"`
BasicAuthPassword string `json:"basic_auth_password,omitempty"`
Remark string `json:"remark,omitempty"`
UpstreamType string `json:"upstream_type,omitempty"`
PagesDeployment *PagesDeployment `json:"pages_deployment,omitempty"`
}
type PagesDeployment struct {
ProjectID uint `json:"project_id"`
ProjectSlug string `json:"project_slug"`
DeploymentID uint `json:"deployment_id"`
DeploymentNumber int `json:"deployment_number"`
Checksum string `json:"checksum"`
EntryFile string `json:"entry_file"`
SPAFallbackEnabled bool `json:"spa_fallback_enabled"`
SPAFallbackPath string `json:"spa_fallback_path"`
APIProxyEnabled bool `json:"api_proxy_enabled"`
APIProxyPath string `json:"api_proxy_path"`
APIProxyPass string `json:"api_proxy_pass"`
APIProxyRewrite string `json:"api_proxy_rewrite"`
LocalRoot string `json:"local_root"`
}
type WAFRuleGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
Enabled bool `json:"enabled"`
IsGlobal bool `json:"is_global"`
BlockStatusCode int `json:"block_status_code"`
BlockResponseBody string `json:"block_response_body,omitempty"`
IPWhitelist []string `json:"ip_whitelist,omitempty"`
IPBlacklist []string `json:"ip_blacklist,omitempty"`
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
CountryWhitelist []string `json:"country_whitelist,omitempty"`
CountryBlacklist []string `json:"country_blacklist,omitempty"`
RegionWhitelist []string `json:"region_whitelist,omitempty"`
RegionBlacklist []string `json:"region_blacklist,omitempty"`
PoWEnabled bool `json:"pow_enabled,omitempty"`
PoWConfig *PoWConfig `json:"pow_config,omitempty"`
}
type WAFIPGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
Enabled bool `json:"enabled"`
IPList []string `json:"ip_list,omitempty"`
}
type WAFBinding struct {
RouteID uint `json:"route_id"`
SiteName string `json:"site_name"`
RuleGroupIDs []uint `json:"rule_group_ids"`
}
type WAFDocument struct {
RuleGroups []WAFRuleGroup `json:"rule_groups"`
IPGroups []WAFIPGroup `json:"ip_groups,omitempty"`
Bindings []WAFBinding `json:"bindings"`
}
type ConfigSnapshot struct {
DefaultServerReturnStatus int `json:"default_server_return_status"`
WorkerProcesses string `json:"worker_processes"`
WorkerConnections int `json:"worker_connections"`
WorkerRlimitNofile int `json:"worker_rlimit_nofile"`
EventsUse string `json:"events_use,omitempty"`
EventsMultiAcceptEnabled bool `json:"events_multi_accept_enabled"`
KeepaliveTimeout int `json:"keepalive_timeout"`
KeepaliveRequests int `json:"keepalive_requests"`
ClientHeaderTimeout int `json:"client_header_timeout"`
ClientBodyTimeout int `json:"client_body_timeout"`
ClientMaxBodySize string `json:"client_max_body_size"`
LargeClientHeaderBuffers string `json:"large_client_header_buffers"`
SendTimeout int `json:"send_timeout"`
ProxyConnectTimeout int `json:"proxy_connect_timeout"`
ProxySendTimeout int `json:"proxy_send_timeout"`
ProxyReadTimeout int `json:"proxy_read_timeout"`
WebsocketEnabled bool `json:"websocket_enabled"`
HTTP3Enabled bool `json:"http3_enabled"`
ProxyRequestBuffering bool `json:"proxy_request_buffering"`
ProxyBufferingEnabled bool `json:"proxy_buffering_enabled"`
ProxyBuffers string `json:"proxy_buffers"`
ProxyBufferSize string `json:"proxy_buffer_size"`
ProxyBusyBuffersSize string `json:"proxy_busy_buffers_size"`
GzipEnabled bool `json:"gzip_enabled"`
GzipMinLength int `json:"gzip_min_length"`
GzipCompLevel int `json:"gzip_comp_level"`
Resolvers string `json:"resolvers,omitempty"`
CacheEnabled bool `json:"cache_enabled"`
CachePath string `json:"cache_path,omitempty"`
CacheLevels string `json:"cache_levels"`
CacheInactive string `json:"cache_inactive"`
CacheMaxSize string `json:"cache_max_size"`
CacheKeyTemplate string `json:"cache_key_template"`
CacheLockEnabled bool `json:"cache_lock_enabled"`
CacheLockTimeout string `json:"cache_lock_timeout"`
CacheUseStale string `json:"cache_use_stale"`
MainConfigTemplate string `json:"main_config_template,omitempty"`
}
type Document struct {
Routes []Route `json:"routes"`
OpenRestyConfig ConfigSnapshot `json:"openresty_config"`
WAF WAFDocument `json:"waf"`
}
type Result struct {
MainConfig string
RouteConfig string
SupportFiles []SupportFile
Checksum string
}
type routeCacheConfig struct {
Enabled bool
Policy string
Rules []string
}
type routeLimitConfig struct {
LimitConnPerServer int
LimitConnPerIP int
LimitRate string
}
type routeUpstreamConfig struct {
Name string
Scheme string
ProxyPassURI string
Servers []string
UsesNamedUpstream bool
}
var requiredMainConfigTemplatePlaceholders = []string{
"{{OpenRestyWorkerProcesses}}",
"{{OpenRestyWorkerConnections}}",
"{{OpenRestyWorkerRlimitNofile}}",
"{{OpenRestyConnectionUpgradeMap}}",
"{{OpenRestyDefaultServerBlock}}",
"{{OpenRestyAccessLogPath}}",
"{{OpenRestyErrorLogPath}}",
"{{OpenRestyEventsUseDirective}}",
"{{OpenRestyEventsMultiAcceptDirective}}",
"{{OpenRestyKeepaliveTimeout}}",
"{{OpenRestyKeepaliveRequests}}",
"{{OpenRestyClientHeaderTimeout}}",
"{{OpenRestyClientBodyTimeout}}",
"{{OpenRestyClientMaxBodySize}}",
"{{OpenRestyLargeClientHeaderBuffers}}",
"{{OpenRestySendTimeout}}",
"{{OpenRestyProxyConnectTimeout}}",
"{{OpenRestyProxySendTimeout}}",
"{{OpenRestyProxyReadTimeout}}",
"{{OpenRestyProxyRequestBuffering}}",
"{{OpenRestyProxyBuffering}}",
"{{OpenRestyProxyBuffers}}",
"{{OpenRestyProxyBufferSize}}",
"{{OpenRestyProxyBusyBuffersSize}}",
"{{OpenRestyGzip}}",
"{{OpenRestyGzipMinLength}}",
"{{OpenRestyGzipCompLevel}}",
"{{OpenRestyCacheBlock}}",
"{{OpenRestyRouteConfigInclude}}",
}
@@ -0,0 +1,14 @@
package security
import "golang.org/x/crypto/bcrypt"
func Password2Hash(password string) (string, error) {
passwordBytes := []byte(password)
hashedPassword, err := bcrypt.GenerateFromPassword(passwordBytes, bcrypt.DefaultCost)
return string(hashedPassword), err
}
func ValidatePasswordAndHash(password string, hash string) bool {
err := bcrypt.CompareHashAndPassword([]byte(hash), []byte(password))
return err == nil
}
+37
View File
@@ -0,0 +1,37 @@
package security
import "crypto/rand"
func GenerateRandomString(length int) string {
if length <= 0 {
return ""
}
const charset = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
const n = byte(len(charset))
const threshold = byte(256 - (256 % len(charset)))
out := make([]byte, 0, length)
buf := make([]byte, length)
for len(out) < length {
if _, err := rand.Read(buf); err != nil {
return ""
}
for _, b := range buf {
if b < threshold {
out = append(out, charset[int(b%n)])
if len(out) == length {
break
}
}
}
}
return string(out)
}
func GeneratePassword() string {
return GenerateRandomString(12)
}
func GenerateToken() string {
return GenerateRandomString(22)
}
@@ -0,0 +1,78 @@
package security
import (
"strings"
"sync"
"time"
"github.com/google/uuid"
)
type verificationValue struct {
code string
time time.Time
}
const (
EmailVerificationPurpose = "v"
PasswordResetPurpose = "r"
)
var verificationMutex sync.Mutex
var verificationMap map[string]verificationValue
var verificationMapMaxSize = 10
var VerificationValidMinutes = 10
func GenerateVerificationCode(length int) string {
code := uuid.New().String()
code = strings.Replace(code, "-", "", -1)
if length == 0 {
return code
}
return code[:length]
}
func RegisterVerificationCodeWithKey(key string, code string, purpose string) {
verificationMutex.Lock()
defer verificationMutex.Unlock()
verificationMap[purpose+key] = verificationValue{
code: code,
time: time.Now(),
}
if len(verificationMap) > verificationMapMaxSize {
removeExpiredPairs()
}
}
func VerifyCodeWithKey(key string, code string, purpose string) bool {
verificationMutex.Lock()
defer verificationMutex.Unlock()
value, okay := verificationMap[purpose+key]
now := time.Now()
if !okay || int(now.Sub(value.time).Seconds()) >= VerificationValidMinutes*60 {
return false
}
return code == value.code
}
func DeleteKey(key string, purpose string) {
verificationMutex.Lock()
defer verificationMutex.Unlock()
delete(verificationMap, purpose+key)
}
// no lock inside, so the caller must lock the verificationMap before calling!
func removeExpiredPairs() {
now := time.Now()
for key := range verificationMap {
if int(now.Sub(verificationMap[key].time).Seconds()) >= VerificationValidMinutes*60 {
delete(verificationMap, key)
}
}
}
func init() {
verificationMutex.Lock()
defer verificationMutex.Unlock()
verificationMap = make(map[string]verificationValue)
}
+76
View File
@@ -0,0 +1,76 @@
package utils
import (
"sort"
"strings"
"time"
)
// Unique returns a new slice containing only the unique elements of the input slice,
// preserving their original order.
func Unique[T comparable](slice []T) []T {
if slice == nil {
return nil
}
seen := make(map[T]struct{})
result := make([]T, 0)
for _, item := range slice {
if _, ok := seen[item]; ok {
continue
}
seen[item] = struct{}{}
result = append(result, item)
}
return result
}
// UniqueAndCleanStringSlice trims spaces, removes empty elements, and returns only the unique elements
// of the input string slice. It preserves order and returns nil if the resulting slice is empty.
func UniqueAndCleanStringSlice(slice []string) []string {
if slice == nil {
return nil
}
seen := make(map[string]struct{})
result := make([]string, 0)
for _, item := range slice {
trimmed := strings.TrimSpace(item)
if trimmed == "" {
continue
}
if _, ok := seen[trimmed]; ok {
continue
}
seen[trimmed] = struct{}{}
result = append(result, trimmed)
}
if len(result) == 0 {
return nil
}
return result
}
// IdentifiableTimeRecord represents a database record that has a unique ID and a primary timestamp field.
type IdentifiableTimeRecord interface {
GetID() uint
GetTime() time.Time
}
// SortAndLimitRecords sorts a slice of IdentifiableTimeRecord descendingly by their timestamp (and ID as a tie-breaker),
// and limits the slice to the specified size if limit > 0.
func SortAndLimitRecords[T IdentifiableTimeRecord](rows []T, limit int) []T {
if len(rows) == 0 {
return rows
}
sort.Slice(rows, func(i, j int) bool {
ti := rows[i].GetTime()
tj := rows[j].GetTime()
if ti.Equal(tj) {
return rows[i].GetID() > rows[j].GetID()
}
return ti.After(tj)
})
if limit > 0 && len(rows) > limit {
rows = rows[:limit]
}
return rows
}
+12
View File
@@ -0,0 +1,12 @@
package utils
import "strings"
// TrimStringFields trims leading and trailing spaces from all provided string pointers.
func TrimStringFields(fields ...*string) {
for _, f := range fields {
if f != nil {
*f = strings.TrimSpace(*f)
}
}
}
+373
View File
@@ -0,0 +1,373 @@
package uptimekuma
import (
"context"
"encoding/json"
"fmt"
"io"
"log/slog"
"net/http"
"strconv"
"strings"
"sync"
"time"
)
type UptimeKumaMonitor struct {
ID int `json:"id"`
Name string `json:"name"`
Url string `json:"url"`
Type string `json:"type"`
Interval int `json:"interval"`
MaxRetries int `json:"maxretries"`
RetryInterval int `json:"retryInterval"`
Timeout int `json:"timeout"`
Tags []UptimeKumaTag `json:"tags"`
}
type UptimeKumaTag struct {
ID int `json:"tag_id"`
Name string `json:"name"`
Color string `json:"color"`
}
type UptimeKumaTagItem struct {
ID int `json:"id"`
Name string `json:"name"`
Color string `json:"color"`
}
type SocketIOClient struct {
baseURL string
httpClient *http.Client
sid string
ackMutex sync.Mutex
ackID int
ackChanMap map[int]chan string
doneChan chan struct{}
closeOnce sync.Once
monitorListMutex sync.RWMutex
monitorList map[string]UptimeKumaMonitor
monitorListChan chan struct{}
monitorListOnce sync.Once
ctx context.Context
cancel context.CancelFunc
err error
}
func NewSocketIOClient(baseURL string) *SocketIOClient {
ctx, cancel := context.WithCancel(context.Background())
return &SocketIOClient{
baseURL: strings.TrimSuffix(baseURL, "/"),
httpClient: &http.Client{
Timeout: 60 * time.Second,
},
ackChanMap: make(map[int]chan string),
doneChan: make(chan struct{}),
monitorListChan: make(chan struct{}),
monitorList: make(map[string]UptimeKumaMonitor),
ctx: ctx,
cancel: cancel,
}
}
func (c *SocketIOClient) Connect() error {
slog.Debug("Uptime Kuma client starting handshake", "baseURL", c.baseURL)
// 1. Handshake
u := fmt.Sprintf("%s/socket.io/?EIO=4&transport=polling", c.baseURL)
reqHandshake, err := http.NewRequestWithContext(c.ctx, "GET", u, nil)
if err != nil {
return fmt.Errorf("create handshake request failed: %w", err)
}
resp, err := c.httpClient.Do(reqHandshake)
if err != nil {
slog.Error("Uptime Kuma handshake connection failed", "url", u, "error", err)
return fmt.Errorf("handshake request failed: %w", err)
}
defer resp.Body.Close()
bs, err := io.ReadAll(resp.Body)
if err != nil {
slog.Error("Failed to read Uptime Kuma handshake response body", "error", err)
return fmt.Errorf("read handshake body failed: %w", err)
}
bodyStr := string(bs)
slog.Debug("Received handshake response from Uptime Kuma", "body", bodyStr)
if len(bodyStr) == 0 || bodyStr[0] != '0' {
return fmt.Errorf("invalid handshake response format: %s", bodyStr)
}
var hs struct {
Sid string `json:"sid"`
}
if err := json.Unmarshal([]byte(bodyStr[1:]), &hs); err != nil {
return fmt.Errorf("unmarshal handshake sid failed: %w", err)
}
c.sid = hs.Sid
slog.Debug("Uptime Kuma handshake success", "sid", c.sid)
// 2. Namespace Connect
slog.Debug("Sending namespace connect request to Uptime Kuma", "sid", c.sid)
connectURL := fmt.Sprintf("%s/socket.io/?EIO=4&transport=polling&sid=%s", c.baseURL, c.sid)
req, err := http.NewRequestWithContext(c.ctx, "POST", connectURL, strings.NewReader("40"))
if err != nil {
return fmt.Errorf("create connect request failed: %w", err)
}
req.Header.Set("Content-Type", "text/plain;charset=UTF-8")
respConnect, err := c.httpClient.Do(req)
if err != nil {
slog.Error("Uptime Kuma namespace connect request failed", "sid", c.sid, "error", err)
return fmt.Errorf("namespace connect failed: %w", err)
}
respConnect.Body.Close()
slog.Debug("Namespace connected successfully to Uptime Kuma", "sid", c.sid)
// 3. Start Polling Loop
go c.pollLoop()
return nil
}
func (c *SocketIOClient) pollLoop() {
slog.Debug("Uptime Kuma polling loop started", "sid", c.sid)
defer c.Close()
for {
select {
case <-c.doneChan:
slog.Debug("Uptime Kuma polling loop stopped (doneChan closed)", "sid", c.sid)
return
default:
}
u := fmt.Sprintf("%s/socket.io/?EIO=4&transport=polling&sid=%s", c.baseURL, c.sid)
reqPoll, err := http.NewRequestWithContext(c.ctx, "GET", u, nil)
if err != nil {
slog.Error("Failed to create Uptime Kuma polling request", "sid", c.sid, "error", err)
c.err = err
return
}
resp, err := c.httpClient.Do(reqPoll)
if err != nil {
slog.Error("Uptime Kuma polling request failed", "sid", c.sid, "error", err)
c.err = err
return
}
bs, err := io.ReadAll(resp.Body)
resp.Body.Close()
if err != nil {
slog.Error("Failed to read Uptime Kuma polling body", "sid", c.sid, "error", err)
c.err = err
return
}
bodyStr := string(bs)
if len(bodyStr) == 0 {
continue
}
slog.Debug("Received polling payload from Uptime Kuma", "length", len(bodyStr))
packets := strings.Split(bodyStr, "\x1e")
for _, pkt := range packets {
if len(pkt) == 0 {
continue
}
engineIOType := pkt[0]
payload := pkt[1:]
slog.Debug("Parsing engine.io packet", "type", string(engineIOType), "payload_len", len(payload))
switch engineIOType {
case '2': // Ping
slog.Debug("Received engine.io ping, responding with pong", "sid", c.sid)
c.sendPong()
case '4': // Message
if len(payload) == 0 {
continue
}
socketIOType := payload[0]
socketIOPayload := payload[1:]
slog.Debug("Parsing socket.io packet", "type", string(socketIOType), "payload", socketIOPayload)
switch socketIOType {
case '2': // Event
c.handleEvent(socketIOPayload)
case '3': // Ack
c.handleAck(socketIOPayload)
}
}
}
}
}
func (c *SocketIOClient) sendPong() {
u := fmt.Sprintf("%s/socket.io/?EIO=4&transport=polling&sid=%s", c.baseURL, c.sid)
req, err := http.NewRequestWithContext(c.ctx, "POST", u, strings.NewReader("3"))
if err != nil {
return
}
req.Header.Set("Content-Type", "text/plain;charset=UTF-8")
resp, err := c.httpClient.Do(req)
if err == nil {
resp.Body.Close()
}
}
func (c *SocketIOClient) handleEvent(payload string) {
var arr []json.RawMessage
if err := json.Unmarshal([]byte(payload), &arr); err != nil || len(arr) < 2 {
return
}
var eventName string
if err := json.Unmarshal(arr[0], &eventName); err != nil {
return
}
if eventName == "monitorList" {
var list map[string]UptimeKumaMonitor
if err := json.Unmarshal(arr[1], &list); err == nil {
c.monitorListMutex.Lock()
c.monitorList = list
c.monitorListMutex.Unlock()
c.monitorListOnce.Do(func() {
close(c.monitorListChan)
})
}
}
}
func (c *SocketIOClient) handleAck(payload string) {
idx := strings.IndexByte(payload, '[')
if idx == -1 {
return
}
ackIDStr := payload[:idx]
ackID, err := strconv.Atoi(ackIDStr)
if err != nil {
return
}
c.ackMutex.Lock()
ch, ok := c.ackChanMap[ackID]
if ok {
delete(c.ackChanMap, ackID)
c.ackMutex.Unlock()
select {
case ch <- payload[idx:]:
default:
}
} else {
c.ackMutex.Unlock()
}
}
func (c *SocketIOClient) Emit(event string, args ...any) (string, error) {
c.ackMutex.Lock()
id := c.ackID
c.ackID++
ch := make(chan string, 1)
c.ackChanMap[id] = ch
c.ackMutex.Unlock()
payloadArr := []any{event}
payloadArr = append(payloadArr, args...)
bs, err := json.Marshal(payloadArr)
if err != nil {
c.ackMutex.Lock()
delete(c.ackChanMap, id)
c.ackMutex.Unlock()
slog.Error("Failed to marshal event payload", "event", event, "error", err)
return "", err
}
body := fmt.Sprintf("42%d%s", id, string(bs))
slog.Debug("Emitting Socket.IO event", "event", event, "ackID", id, "payload", string(bs))
u := fmt.Sprintf("%s/socket.io/?EIO=4&transport=polling&sid=%s", c.baseURL, c.sid)
req, err := http.NewRequestWithContext(c.ctx, "POST", u, strings.NewReader(body))
if err != nil {
c.ackMutex.Lock()
delete(c.ackChanMap, id)
c.ackMutex.Unlock()
return "", err
}
req.Header.Set("Content-Type", "text/plain;charset=UTF-8")
resp, err := c.httpClient.Do(req)
if err != nil {
c.ackMutex.Lock()
delete(c.ackChanMap, id)
c.ackMutex.Unlock()
slog.Error("Failed to send Emit request", "event", event, "ackID", id, "error", err)
return "", err
}
resp.Body.Close()
select {
case result := <-ch:
slog.Debug("Received Ack for event", "event", event, "ackID", id, "response", result)
return result, nil
case <-time.After(10 * time.Second):
c.ackMutex.Lock()
delete(c.ackChanMap, id)
c.ackMutex.Unlock()
slog.Error("Timeout waiting for event Ack", "event", event, "ackID", id)
return "", fmt.Errorf("timeout waiting for ack for event: %s", event)
case <-c.doneChan:
c.ackMutex.Lock()
delete(c.ackChanMap, id)
c.ackMutex.Unlock()
slog.Error("Client closed while waiting for event Ack", "event", event, "ackID", id)
return "", fmt.Errorf("client closed while waiting for event ack: %s", event)
}
}
func (c *SocketIOClient) Close() {
c.closeOnce.Do(func() {
c.cancel()
close(c.doneChan)
})
}
func (c *SocketIOClient) GetMonitorListChan() <-chan struct{} {
return c.monitorListChan
}
func (c *SocketIOClient) GetMonitorList() map[string]UptimeKumaMonitor {
c.monitorListMutex.RLock()
defer c.monitorListMutex.RUnlock()
// Return a copy to prevent concurrent map read/write access
m := make(map[string]UptimeKumaMonitor, len(c.monitorList))
for k, v := range c.monitorList {
m[k] = v
}
return m
}
func ParseAckResponse(response string, target any) error {
var arr []json.RawMessage
if err := json.Unmarshal([]byte(response), &arr); err != nil || len(arr) == 0 {
return fmt.Errorf("invalid ack response format: %s", response)
}
var status struct {
Ok bool `json:"ok"`
Msg string `json:"msg"`
}
if err := json.Unmarshal(arr[0], &status); err == nil {
if !status.Ok {
errMsg := status.Msg
if errMsg == "" {
errMsg = "unknown error from Uptime Kuma"
}
return fmt.Errorf("Uptime Kuma error response: %s", errMsg)
}
}
if target != nil {
return json.Unmarshal(arr[0], target)
}
return nil
}
@@ -0,0 +1,9 @@
package validation
import "github.com/go-playground/validator/v10"
var Validate *validator.Validate
func init() {
Validate = validator.New()
}
+15
View File
@@ -0,0 +1,15 @@
package utils
import "fmt"
func Interface2String(inter interface{}) string {
switch inter.(type) {
case string:
return inter.(string)
case int:
return fmt.Sprintf("%d", inter.(int))
case float64:
return fmt.Sprintf("%f", inter.(float64))
}
return "Not Implemented"
}
+208
View File
@@ -0,0 +1,208 @@
package utils
import (
"strconv"
"strings"
)
type VersionInfo struct {
Valid bool
IsDev bool
Numbers []int
Prerelease []string
GitDescribeDistance int
GitDescribeTail []string
}
func ParseVersionInfo(version string) VersionInfo {
normalized := strings.TrimSpace(strings.TrimPrefix(version, "v"))
if normalized == "" || normalized == "dev" {
return VersionInfo{IsDev: strings.EqualFold(normalized, "dev")}
}
base := normalized
prerelease := ""
if separator := strings.IndexRune(normalized, '-'); separator >= 0 {
base = normalized[:separator]
prerelease = normalized[separator+1:]
}
segments := strings.Split(base, ".")
parts := make([]int, 0, len(segments))
for _, segment := range segments {
segment = strings.TrimSpace(segment)
if segment == "" {
parts = append(parts, 0)
continue
}
numeric := strings.Builder{}
for _, r := range segment {
if r < '0' || r > '9' {
break
}
numeric.WriteRune(r)
}
if numeric.Len() == 0 {
parts = append(parts, 0)
continue
}
value, err := strconv.Atoi(numeric.String())
if err != nil {
return VersionInfo{}
}
parts = append(parts, value)
}
info := VersionInfo{Valid: len(parts) > 0, Numbers: parts}
if prerelease != "" {
identifiers := splitPrereleaseIdentifiers(prerelease)
if distance, tail, ok := parseGitDescribeIdentifiers(identifiers); ok {
info.GitDescribeDistance = distance
info.GitDescribeTail = tail
} else {
info.Prerelease = identifiers
}
}
return info
}
func parseGitDescribeIdentifiers(identifiers []string) (int, []string, bool) {
if len(identifiers) < 2 {
return 0, nil, false
}
distance, err := strconv.Atoi(strings.TrimSpace(identifiers[0]))
if err != nil || distance <= 0 {
return 0, nil, false
}
commitToken := strings.TrimSpace(identifiers[1])
if commitToken == "" || !strings.HasPrefix(strings.ToLower(commitToken), "g") {
return 0, nil, false
}
return distance, identifiers[1:], true
}
func splitPrereleaseIdentifiers(value string) []string {
parts := strings.FieldsFunc(strings.TrimSpace(value), func(r rune) bool {
return r == '.' || r == '-'
})
filtered := make([]string, 0, len(parts))
for _, part := range parts {
part = strings.TrimSpace(part)
if part != "" {
filtered = append(filtered, part)
}
}
return filtered
}
// CompareVersions compares two version strings.
// Returns -1 if left < right, 1 if left > right, and 0 if they are equal.
func CompareVersions(local, remote string) int {
left := ParseVersionInfo(local)
right := ParseVersionInfo(remote)
if left.IsDev {
if right.Valid {
return -1
}
return 0
}
if !left.Valid || !right.Valid {
return 0
}
maxLen := len(left.Numbers)
if len(right.Numbers) > maxLen {
maxLen = len(right.Numbers)
}
for index := 0; index < maxLen; index++ {
leftValue := 0
rightValue := 0
if index < len(left.Numbers) {
leftValue = left.Numbers[index]
}
if index < len(right.Numbers) {
rightValue = right.Numbers[index]
}
if leftValue < rightValue {
return -1
}
if leftValue > rightValue {
return 1
}
}
if left.GitDescribeDistance != right.GitDescribeDistance {
if left.GitDescribeDistance < right.GitDescribeDistance {
return -1
}
return 1
}
if left.GitDescribeDistance > 0 || right.GitDescribeDistance > 0 {
maxLen = len(left.GitDescribeTail)
if len(right.GitDescribeTail) > maxLen {
maxLen = len(right.GitDescribeTail)
}
for index := 0; index < maxLen; index++ {
if index >= len(left.GitDescribeTail) {
return -1
}
if index >= len(right.GitDescribeTail) {
return 1
}
if left.GitDescribeTail[index] < right.GitDescribeTail[index] {
return -1
}
if left.GitDescribeTail[index] > right.GitDescribeTail[index] {
return 1
}
}
return 0
}
if len(left.Prerelease) == 0 && len(right.Prerelease) == 0 {
return 0
}
if len(left.Prerelease) == 0 {
return 1
}
if len(right.Prerelease) == 0 {
return -1
}
maxLen = len(left.Prerelease)
if len(right.Prerelease) > maxLen {
maxLen = len(right.Prerelease)
}
for index := 0; index < maxLen; index++ {
if index >= len(left.Prerelease) {
return -1
}
if index >= len(right.Prerelease) {
return 1
}
leftPart := left.Prerelease[index]
rightPart := right.Prerelease[index]
leftNumber, leftErr := strconv.Atoi(leftPart)
rightNumber, rightErr := strconv.Atoi(rightPart)
switch {
case leftErr == nil && rightErr == nil:
if leftNumber < rightNumber {
return -1
}
if leftNumber > rightNumber {
return 1
}
case leftErr == nil:
return -1
case rightErr == nil:
return 1
default:
if leftPart < rightPart {
return -1
}
if leftPart > rightPart {
return 1
}
}
}
return 0
}
+222
View File
@@ -0,0 +1,222 @@
package wsclient
import (
"context"
"encoding/json"
"errors"
"log/slog"
"net"
"net/http"
"net/url"
"strings"
"time"
"golang.org/x/net/websocket"
)
type Config struct {
BaseURL string
Token string
Timeout time.Duration
HeaderKey string // e.g. "X-Agent-Token", "X-Tunnel-Token"
WSPath string // e.g. "/api/relay/ws", "/api/agent/ws", "/api/flared/ws"
}
type Client struct {
cfg Config
}
type WSMessage struct {
Type string `json:"type"`
Payload json.RawMessage `json:"payload,omitempty"`
}
type MessageHandler interface {
OnConnect(ctx context.Context) error
HandleMessage(ctx context.Context, msg WSMessage) error
OnClose(err error)
}
type Connection struct {
Conn *websocket.Conn
URL string
ReadTimeout time.Duration
}
func New(cfg Config) *Client {
cfg.BaseURL = strings.TrimRight(cfg.BaseURL, "/")
cfg.Token = strings.TrimSpace(cfg.Token)
cfg.HeaderKey = strings.TrimSpace(cfg.HeaderKey)
cfg.WSPath = strings.TrimSpace(cfg.WSPath)
return &Client{
cfg: cfg,
}
}
func (c *Client) SetToken(token string) {
c.cfg.Token = strings.TrimSpace(token)
}
func (c *Client) URL() string {
wsURL, err := c.BuildWebsocketURL()
if err != nil {
return ""
}
return wsURL
}
func (c *Client) BuildWebsocketURL() (string, error) {
parsed, err := url.Parse(c.cfg.BaseURL)
if err != nil {
return "", err
}
switch parsed.Scheme {
case "http":
parsed.Scheme = "ws"
case "https":
parsed.Scheme = "wss"
case "ws", "wss":
default:
return "", errors.New("server_url scheme must be http, https, ws, or wss")
}
wsPath := c.cfg.WSPath
if !strings.HasPrefix(wsPath, "/") {
wsPath = "/" + wsPath
}
parsed.Path = strings.TrimRight(parsed.Path, "/") + wsPath
parsed.RawQuery = ""
parsed.Fragment = ""
return parsed.String(), nil
}
func (c *Client) Connect(ctx context.Context) (*Connection, error) {
wsURL, err := c.BuildWebsocketURL()
if err != nil {
return nil, err
}
if c.cfg.Token == "" {
return nil, errors.New("ws token is empty")
}
origin := c.cfg.BaseURL
if origin == "" {
origin = "http://localhost"
}
config, err := websocket.NewConfig(wsURL, origin)
if err != nil {
return nil, err
}
config.Header = http.Header{}
if c.cfg.HeaderKey != "" {
config.Header.Set(c.cfg.HeaderKey, c.cfg.Token)
}
if c.cfg.Timeout > 0 {
config.Dialer = &net.Dialer{Timeout: c.cfg.Timeout}
}
slog.Debug("ws dialing server", "url", wsURL)
conn, err := config.DialContext(ctx)
if err != nil {
return nil, err
}
slog.Debug("ws dial succeeded", "url", wsURL)
return &Connection{Conn: conn, URL: wsURL, ReadTimeout: websocketReadTimeout(c.cfg.Timeout)}, nil
}
func (conn *Connection) SendMessage(msgType string, payload any) error {
if conn == nil || conn.Conn == nil {
return errors.New("ws connection is nil")
}
slog.Debug("ws sending message", "type", msgType)
// Create the outbound message wrapper
message := struct {
Type string `json:"type"`
Payload any `json:"payload,omitempty"`
}{
Type: msgType,
Payload: payload,
}
_ = conn.Conn.SetWriteDeadline(time.Now().Add(5 * time.Second))
return websocket.JSON.Send(conn.Conn, message)
}
func (conn *Connection) Receive(target any) error {
if conn == nil || conn.Conn == nil {
return errors.New("ws connection is nil")
}
if conn.ReadTimeout > 0 {
_ = conn.Conn.SetReadDeadline(time.Now().Add(conn.ReadTimeout))
}
err := websocket.JSON.Receive(conn.Conn, target)
if err != nil {
var netErr net.Error
if errors.As(err, &netErr) && netErr.Timeout() {
slog.Debug("ws receive timeout waiting for server message", "timeout", conn.ReadTimeout)
}
return err
}
return nil
}
func websocketReadTimeout(requestTimeout time.Duration) time.Duration {
timeout := requestTimeout * 6
if timeout < 75*time.Second {
return 75 * time.Second
}
return timeout
}
func (conn *Connection) RunReceiveLoop(ctx context.Context, handler MessageHandler) error {
doneChan := make(chan struct{})
defer close(doneChan)
go func() {
select {
case <-ctx.Done():
_ = conn.Close()
case <-doneChan:
}
}()
if err := handler.OnConnect(ctx); err != nil {
handler.OnClose(err)
return err
}
for {
select {
case <-ctx.Done():
return ctx.Err()
default:
}
var raw WSMessage
if err := conn.Receive(&raw); err != nil {
handler.OnClose(err)
return err
}
switch raw.Type {
case "ping":
slog.Debug("ws received ping from server, replying with pong")
if err := conn.SendMessage("pong", nil); err != nil {
slog.Error("ws send pong response failed", "error", err)
}
case "pong":
slog.Debug("ws received pong response from server")
default:
if err := handler.HandleMessage(ctx, raw); err != nil {
slog.Error("ws handler failed to process message", "type", raw.Type, "error", err)
return err
}
}
}
}
func (conn *Connection) Close() error {
if conn == nil || conn.Conn == nil {
return nil
}
return conn.Conn.Close()
}