mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-30 06:16:37 +08:00
[优化] go 引用调整
This commit is contained in:
@@ -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 = ®istration.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
|
||||
}
|
||||
@@ -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),
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
Reference in New Issue
Block a user