[优化] go 引用调整

This commit is contained in:
ryan
2026-06-06 10:26:20 +08:00
parent ee1110b752
commit 3cfefb4367
552 changed files with 1642 additions and 2185 deletions
@@ -0,0 +1,25 @@
package geoip
import (
"fmt"
"net"
)
type EmptyProvider struct{}
func (e *EmptyProvider) Name() string {
return "EmptyProvider"
}
func (e *EmptyProvider) Initialize() error {
return nil
}
func (e *EmptyProvider) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
return nil, fmt.Errorf("you are using an empty GeoIP provider, please set a valid provider")
}
func (e *EmptyProvider) UpdateDatabase() error {
return fmt.Errorf("you are using an empty GeoIP provider, please set a valid provider")
}
func (e *EmptyProvider) Close() error {
return nil
}
+242
View File
@@ -0,0 +1,242 @@
package geoip
import (
"fmt"
"log/slog"
"net"
"strings"
"sync"
"time"
"unicode"
ristretto "github.com/dgraph-io/ristretto/v2"
)
var CurrentProvider GeoIPService
var geoCache *providerCache
var providerMutex sync.RWMutex
var providerFactory = newProvider
const (
ProviderDisabled = "disabled"
ProviderMaxMind = "mmdb"
ProviderIPAPI = "ip-api"
ProviderGeoJS = "geojs"
ProviderIPInfo = "ipinfo"
)
type GeoInfo struct {
ISOCode string
Name string
Latitude *float64
Longitude *float64
}
func init() {
CurrentProvider = &EmptyProvider{}
geoCache = newProviderCache(48 * time.Hour)
}
// GeoIPService 接口定义了获取地理位置信息的核心方法。
type GeoIPService interface {
Name() string
GetGeoInfo(ip net.IP) (*GeoInfo, error)
UpdateDatabase() error
Close() error
}
type cachedGeoInfo struct {
info *GeoInfo
expiresAt time.Time
}
type providerCache struct {
items *ristretto.Cache[string, cachedGeoInfo]
duration time.Duration
}
func newProviderCache(duration time.Duration) *providerCache {
items, err := ristretto.NewCache(&ristretto.Config[string, cachedGeoInfo]{
NumCounters: 1e5,
MaxCost: 2e4,
BufferItems: 64,
})
if err != nil {
panic(err)
}
return &providerCache{
items: items,
duration: duration,
}
}
func (c *providerCache) Get(key string) (*GeoInfo, bool) {
entry, ok := c.items.Get(key)
if !ok {
return nil, false
}
if time.Now().After(entry.expiresAt) {
c.items.Del(key)
return nil, false
}
return entry.info, true
}
func (c *providerCache) Set(key string, info *GeoInfo) {
c.items.Set(key, cachedGeoInfo{
info: info,
expiresAt: time.Now().Add(c.duration),
}, 1)
c.items.Wait()
}
func (c *providerCache) Flush() {
c.items.Clear()
}
func GetRegionUnicodeEmoji(isoCode string) string {
if len(isoCode) != 2 {
return ""
}
isoCode = strings.ToUpper(isoCode)
if !unicode.IsLetter(rune(isoCode[0])) || !unicode.IsLetter(rune(isoCode[1])) {
return ""
}
rune1 := rune(0x1F1E6 + (rune(isoCode[0]) - 'A'))
rune2 := rune(0x1F1E6 + (rune(isoCode[1]) - 'A'))
return string(rune1) + string(rune2)
}
func InitGeoIP(provider string) {
providerName := normalizeProvider(provider)
nextProvider, err := providerFactory(providerName)
if err != nil {
slog.Error("initialize GeoIP provider failed", "provider", providerName, "error", err)
nextProvider = &EmptyProvider{}
}
setProvider(nextProvider)
if providerName == ProviderDisabled {
slog.Info("GeoIP provider disabled")
return
}
slog.Info("GeoIP provider configured", "provider", CurrentProvider.Name())
}
func GetGeoInfo(ip net.IP) (*GeoInfo, error) {
if ip == nil {
return nil, fmt.Errorf("IP address cannot be nil")
}
provider := getProvider()
cacheKey := provider.Name() + ":" + ip.String()
if cachedInfo, found := geoCache.Get(cacheKey); found {
return cachedInfo, nil
}
info, err := provider.GetGeoInfo(ip)
if err == nil && info != nil {
geoCache.Set(cacheKey, info)
}
return info, err
}
func LookupGeoInfoWithProvider(providerName string, ip net.IP) (*GeoInfo, error) {
if ip == nil {
return nil, fmt.Errorf("IP address cannot be nil")
}
provider, err := providerFactory(normalizeProvider(providerName))
if err != nil {
return nil, err
}
defer func() {
if closeErr := provider.Close(); closeErr != nil {
slog.Warn("close temporary GeoIP provider failed", "provider", provider.Name(), "error", closeErr)
}
}()
return provider.GetGeoInfo(ip)
}
func UpdateDatabase() error {
err := getProvider().UpdateDatabase()
if err == nil {
geoCache.Flush()
slog.Info("GeoIP cache cleared due to database update.")
}
return err
}
func IsValidProvider(provider string) bool {
switch normalizeProvider(provider) {
case ProviderDisabled, ProviderMaxMind, ProviderIPAPI, ProviderGeoJS, ProviderIPInfo:
return true
default:
return false
}
}
func normalizeProvider(provider string) string {
normalized := strings.TrimSpace(strings.ToLower(provider))
if normalized == "" {
return ProviderDisabled
}
return normalized
}
func newProvider(provider string) (GeoIPService, error) {
switch provider {
case ProviderDisabled:
return &EmptyProvider{}, nil
case ProviderMaxMind:
return NewMaxMindGeoIPService()
case ProviderIPAPI:
return NewIPAPIService()
case ProviderGeoJS:
return NewGeoJSService()
case ProviderIPInfo:
return NewIPInfoService()
default:
return nil, fmt.Errorf("unsupported GeoIP provider %q", provider)
}
}
func setProvider(provider GeoIPService) {
providerMutex.Lock()
previous := CurrentProvider
CurrentProvider = provider
providerMutex.Unlock()
geoCache.Flush()
if previous != nil && previous != provider {
if err := previous.Close(); err != nil {
slog.Warn("close previous GeoIP provider failed", "error", err)
}
}
}
func getProvider() GeoIPService {
providerMutex.RLock()
defer providerMutex.RUnlock()
if CurrentProvider == nil {
return &EmptyProvider{}
}
return CurrentProvider
}
func float64Pointer(value float64) *float64 {
return &value
}
func ProviderFactoryForTest() func(string) (GeoIPService, error) {
return providerFactory
}
func SetProviderFactoryForTest(factory func(string) (GeoIPService, error)) {
if factory == nil {
providerFactory = newProvider
return
}
providerFactory = factory
}
@@ -0,0 +1,99 @@
package geoip
import (
"net"
"testing"
)
type fakeProvider struct {
calls int
}
func (f *fakeProvider) Name() string {
return "fake"
}
func (f *fakeProvider) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
f.calls++
return &GeoInfo{
ISOCode: "CN",
Name: "China",
}, nil
}
func (f *fakeProvider) UpdateDatabase() error {
return nil
}
func (f *fakeProvider) Close() error {
return nil
}
func TestGetGeoInfoCachesByProviderAndIP(t *testing.T) {
originalProvider := CurrentProvider
geoCache.Flush()
fake := &fakeProvider{}
CurrentProvider = fake
defer func() {
CurrentProvider = originalProvider
}()
ip := net.ParseIP("8.8.8.8")
record, err := GetGeoInfo(ip)
if err != nil {
t.Fatalf("expected nil error, got %v", err)
}
if record == nil || record.ISOCode != "CN" {
t.Fatalf("expected cached record, got %#v", record)
}
_, err = GetGeoInfo(ip)
if err != nil {
t.Fatalf("expected nil error on second call, got %v", err)
}
if fake.calls != 1 {
t.Fatalf("expected provider to be called once, got %d", fake.calls)
}
}
func TestUnicodeEmoji(t *testing.T) {
emoji := GetRegionUnicodeEmoji("CN")
if emoji != "🇨🇳" {
t.Errorf("expected emoji for CN, got %s", emoji)
}
}
func TestIsValidProvider(t *testing.T) {
cases := map[string]bool{
"disabled": true,
"mmdb": true,
"ip-api": true,
"geojs": true,
"ipinfo": true,
"unknown": false,
}
for provider, want := range cases {
if got := IsValidProvider(provider); got != want {
t.Fatalf("provider %s validity mismatch: want %v, got %v", provider, want, got)
}
}
}
func TestLookupGeoInfoWithProviderUsesTemporaryProvider(t *testing.T) {
previousFactory := providerFactory
providerFactory = func(provider string) (GeoIPService, error) {
return &fakeProvider{}, nil
}
defer func() {
providerFactory = previousFactory
}()
info, err := LookupGeoInfoWithProvider("ipinfo", net.ParseIP("8.8.8.8"))
if err != nil {
t.Fatalf("expected lookup to succeed, got %v", err)
}
if info == nil || info.ISOCode != "CN" || info.Name != "China" {
t.Fatalf("unexpected geo info: %#v", info)
}
}
+84
View File
@@ -0,0 +1,84 @@
package geoip
import (
"encoding/json"
"fmt"
"net"
"net/http"
"time"
)
// GeoJSService 使用 geojs.io 服务实现 GeoIPService 接口。
type GeoJSService struct {
Client *http.Client
}
// geoJSResponse 定义了 geojs.io 服务返回的 JSON 响应的结构。
// 我们只定义我们需要的字段。
type geoJSResponse struct {
Country string `json:"country"`
CountryCode string `json:"country_code"`
Latitude float64 `json:"latitude,string"`
Longitude float64 `json:"longitude,string"`
// 可以根据需要添加其他字段,例如:
// City string `json:"city"`
// Region string `json:"region"`
}
// NewGeoJSService 创建并返回一个 GeoJSService 的新实例。
func NewGeoJSService() (*GeoJSService, error) {
return &GeoJSService{
Client: &http.Client{
Timeout: 5 * time.Second, // 设置一个合理的超时时间
},
}, nil
}
// Name 返回服务的名称。
func (s *GeoJSService) Name() string {
return "geojs.io"
}
// GetGeoInfo 使用 geojs.io 服务检索给定 IP 地址的地理位置信息。
func (s *GeoJSService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
// GeoJS 的 API 端点
apiURL := fmt.Sprintf("https://get.geojs.io/v1/ip/geo/%s.json", ip.String())
resp, err := s.Client.Get(apiURL)
if err != nil {
return nil, fmt.Errorf("failed to get geo info from geojs.io: %w", err)
}
defer resp.Body.Close()
// 检查响应状态码
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("geojs.io returned non-200 status code: %d", resp.StatusCode)
}
var apiResp geoJSResponse
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
return nil, fmt.Errorf("failed to decode geojs.io response: %w", err)
}
// 检查国家代码是否为空,因为 geojs 对无效/私有IP可能返回200 OK但内容为空
if apiResp.CountryCode == "" {
return nil, fmt.Errorf("geojs.io returned empty geo info for ip: %s", ip.String())
}
return &GeoInfo{
ISOCode: apiResp.CountryCode,
Name: apiResp.Country,
Latitude: float64Pointer(apiResp.Latitude),
Longitude: float64Pointer(apiResp.Longitude),
}, nil
}
// UpdateDatabase 对于 geojs.io 是一个空操作,因为它是一个 Web 服务。
func (s *GeoJSService) UpdateDatabase() error {
return nil
}
// Close 对于 geojs.io 是一个空操作。
func (s *GeoJSService) Close() error {
return nil
}
+86
View File
@@ -0,0 +1,86 @@
package geoip
import (
"encoding/json"
"fmt"
"net"
"net/http"
"time"
)
// IPAPIService 使用 ip-api.com 服务实现 GeoIPService 接口。
type IPAPIService struct {
Client *http.Client
}
// ipAPIResponse 定义了 ip-api.com 服务返回的 JSON 响应的结构。
type ipAPIResponse struct {
Status string `json:"status"`
Message string `json:"message"` // 当 status 为 fail 时出现
Country string `json:"country"`
CountryCode string `json:"countryCode"`
Region string `json:"region"`
RegionName string `json:"regionName"`
City string `json:"city"`
Zip string `json:"zip"`
Lat float64 `json:"lat"`
Lon float64 `json:"lon"`
Timezone string `json:"timezone"`
ISP string `json:"isp"`
Org string `json:"org"`
As string `json:"as"`
Query string `json:"query"`
}
func (s *IPAPIService) Name() string {
return "ip-api.com"
}
// NewIPAPIService 创建并返回一个 IPAPIService 的新实例。
func NewIPAPIService() (*IPAPIService, error) {
return &IPAPIService{
Client: &http.Client{
Timeout: 5 * time.Second, // 设置请求超时
},
}, nil
}
// GetGeoInfo 使用 ip-api.com 服务检索给定 IP 地址的地理位置信息。
func (s *IPAPIService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
// API URL, 使用 fields 参数来仅请求需要的字段
apiURL := fmt.Sprintf("http://ip-api.com/json/%s?fields=status,message,country,countryCode", ip.String())
resp, err := s.Client.Get(apiURL)
if err != nil {
return nil, fmt.Errorf("failed to get geo info from ip-api.com: %w", err)
}
defer resp.Body.Close()
var apiResp ipAPIResponse
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
return nil, fmt.Errorf("failed to decode ip-api.com response: %w", err)
}
if apiResp.Status != "success" {
return nil, fmt.Errorf("ip-api.com returned an error: %s", apiResp.Message)
}
return &GeoInfo{
ISOCode: apiResp.CountryCode,
Name: apiResp.Country,
Latitude: float64Pointer(apiResp.Lat),
Longitude: float64Pointer(apiResp.Lon),
}, nil
}
// UpdateDatabase 对于 ip-api.com 是一个空操作,因为它是一个 Web 服务。
func (s *IPAPIService) UpdateDatabase() error {
// 无需执行任何操作,因为数据由外部服务提供
return nil
}
// Close 对于 ip-api.com 是一个空操作,因为没有需要关闭的持久连接。
func (s *IPAPIService) Close() error {
// 无需执行任何操作
return nil
}
+112
View File
@@ -0,0 +1,112 @@
package geoip
import (
"encoding/json"
"fmt"
"net"
"net/http"
"strconv"
"strings"
"time"
)
// IPInfoService 使用 ipinfo.io 服务实现 GeoIPService 接口。
type IPInfoService struct {
Client *http.Client
// 每天 1000 次请求,限制由 IP 地址的所有人共享。
// APIToken string
}
// ipInfoResponse 定义了 ipinfo.io 服务返回的 JSON 响应的结构,只包含免费额度可用的字段。
type ipInfoResponse struct {
IP string `json:"ip"`
Hostname string `json:"hostname"`
City string `json:"city"`
Region string `json:"region"`
Country string `json:"country"`
CountryCode string `json:"countryCode"` // ipinfo.io 返回 "country" 的 ISO 代码,这里为了与 GeoInfo 保持一致,额外添加一个 CountryCode
Loc string `json:"loc"` // Latitude,Longitude
Org string `json:"org"`
Postal string `json:"postal"`
Timezone string `json:"timezone"`
}
// NewIPInfoService 创建并返回一个 IPInfoService 的新实例。
func NewIPInfoService() (*IPInfoService, error) {
return &IPInfoService{
Client: &http.Client{
Timeout: 5 * time.Second,
},
}, nil
}
// Name 返回服务的名称。
func (s *IPInfoService) Name() string {
return "ipinfo.io"
}
// GetGeoInfo 使用 ipinfo.io 服务检索给定 IP 地址的地理位置信息。
// 免费额度主要提供国家信息。
func (s *IPInfoService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
// IPinfo 免费额度不需要 API token 就可以查询基本的 IP 信息。
// API URL: https://ipinfo.io/json (查询自身IP) 或 https://ipinfo.io/YOUR_IP/json
apiURL := fmt.Sprintf("https://ipinfo.io/%s/json", ip.String())
resp, err := s.Client.Get(apiURL)
if err != nil {
return nil, fmt.Errorf("failed to get geo info from ipinfo.io: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("ipinfo.io returned non-200 status: %d %s", resp.StatusCode, resp.Status)
}
var apiResp ipInfoResponse
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
return nil, fmt.Errorf("failed to decode ipinfo.io response: %w", err)
}
latitude, longitude := parseIPInfoCoordinates(apiResp.Loc)
// IPinfo 的 "country" 字段直接返回 ISO 2-letter code,例如 "US", "CN"
// 我们需要将 "country" 字段作为 ISOCode,并尝试获取其对应的国家名称。
// IPinfo 响应中通常不直接提供完整的国家名称,但我们可以通过 CountryCode 映射。
// 为了简化并符合 GeoInfo 结构,我们直接使用 Country 作为 ISOCode,并尝试从 CountryCode 获取名称。
// 实际上,IPinfo 的 'country' 字段就是 ISO 2-letter code。
// 如果需要完整的国家名称,可能需要一个本地的 ISO 代码到名称的映射。
// 为了与 GetRegionUnicodeEmoji 函数兼容,我们直接使用 country 作为 ISOCode。
return &GeoInfo{
ISOCode: apiResp.Country,
Name: apiResp.Country,
Latitude: latitude,
Longitude: longitude,
}, nil
}
// UpdateDatabase 对于 ipinfo.io 是一个空操作,因为它是一个 Web 服务。
func (s *IPInfoService) UpdateDatabase() error {
// 无需执行任何操作,因为数据由外部服务提供
return nil
}
// Close 对于 ipinfo.io 是一个空操作,因为没有需要关闭的持久连接。
func (s *IPInfoService) Close() error {
// 无需执行任何操作
return nil
}
func parseIPInfoCoordinates(value string) (*float64, *float64) {
parts := strings.Split(strings.TrimSpace(value), ",")
if len(parts) != 2 {
return nil, nil
}
latitudeValue, latErr := strconv.ParseFloat(strings.TrimSpace(parts[0]), 64)
longitudeValue, lonErr := strconv.ParseFloat(strings.TrimSpace(parts[1]), 64)
if latErr != nil || lonErr != nil {
return nil, nil
}
return float64Pointer(latitudeValue), float64Pointer(longitudeValue)
}
@@ -0,0 +1,66 @@
package iputil
import (
"net"
"strings"
)
func NormalizeIP(raw string) string {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return ""
}
ip := net.ParseIP(trimmed)
if ip == nil {
return ""
}
if ipv4 := ip.To4(); ipv4 != nil {
return ipv4.String()
}
return ip.String()
}
func NormalizeRemoteAddr(remoteAddr string) string {
trimmed := strings.TrimSpace(remoteAddr)
if trimmed == "" {
return ""
}
if host, _, err := net.SplitHostPort(trimmed); err == nil {
return NormalizeIP(host)
}
return NormalizeIP(trimmed)
}
func IsPublic(ip net.IP) bool {
if ip == nil {
return false
}
if ipv4 := ip.To4(); ipv4 != nil {
ip = ipv4
}
if !ip.IsGlobalUnicast() || ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsMulticast() || ip.IsUnspecified() {
return false
}
return true
}
func IsPublicString(raw string) bool {
ip := net.ParseIP(strings.TrimSpace(raw))
return IsPublic(ip)
}
func Score(ip net.IP) int {
if ip == nil {
return -1
}
if ipv4 := ip.To4(); ipv4 != nil {
ip = ipv4
}
if !ip.IsGlobalUnicast() || ip.IsLoopback() || ip.IsMulticast() || ip.IsUnspecified() {
return -1
}
if IsPublic(ip) {
return 2
}
return 1
}
@@ -0,0 +1,45 @@
package iputil
import (
"net"
"testing"
)
func TestNormalizeIP(t *testing.T) {
if got := NormalizeIP(" 8.8.8.8 "); got != "8.8.8.8" {
t.Fatalf("unexpected normalized ipv4: %q", got)
}
if got := NormalizeIP("[::1]"); got != "" {
t.Fatalf("expected invalid bracketed host to be rejected, got %q", got)
}
}
func TestNormalizeRemoteAddr(t *testing.T) {
if got := NormalizeRemoteAddr("203.0.113.10:8443"); got != "203.0.113.10" {
t.Fatalf("unexpected remote addr normalization: %q", got)
}
}
func TestIsPublic(t *testing.T) {
if !IsPublic(net.ParseIP("8.8.8.8")) {
t.Fatal("expected public ip to be detected")
}
if IsPublic(net.ParseIP("10.0.0.8")) {
t.Fatal("expected private ip to be rejected")
}
if IsPublic(net.ParseIP("127.0.0.1")) {
t.Fatal("expected loopback ip to be rejected")
}
}
func TestScore(t *testing.T) {
if got := Score(net.ParseIP("8.8.8.8")); got != 2 {
t.Fatalf("unexpected score for public ip: %d", got)
}
if got := Score(net.ParseIP("10.0.0.8")); got != 1 {
t.Fatalf("unexpected score for private ip: %d", got)
}
if got := Score(net.ParseIP("127.0.0.1")); got != -1 {
t.Fatalf("unexpected score for loopback ip: %d", got)
}
}
+171
View File
@@ -0,0 +1,171 @@
package geoip
import (
"fmt"
"io"
"net"
"net/http"
"os"
"path/filepath"
"sync"
"github.com/oschwald/maxminddb-golang"
)
var GeoIpUrl = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb"
var GeoIpFilePath = "./data/GeoLite2-Country.mmdb"
type GeoIpRecord struct {
Country struct {
ISOCode string `maxminddb:"iso_code"`
Names map[string]string `maxminddb:"names"`
} `maxminddb:"country"`
}
type MaxMindGeoIPService struct {
maxMindDBReader *maxminddb.Reader
dbFilePath string
mu sync.RWMutex
}
func (s *MaxMindGeoIPService) Name() string {
return "MaxMind"
}
func NewMaxMindGeoIPService() (*MaxMindGeoIPService, error) {
return NewMaxMindGeoIPServiceWithConfig(GeoIpFilePath, GeoIpUrl)
}
func NewMaxMindGeoIPServiceWithConfig(dbFilePath string, downloadURL string) (*MaxMindGeoIPService, error) {
if dbFilePath == "" {
dbFilePath = GeoIpFilePath
}
if downloadURL == "" {
downloadURL = GeoIpUrl
}
service := &MaxMindGeoIPService{
dbFilePath: dbFilePath,
}
if err := os.MkdirAll(filepath.Dir(service.dbFilePath), os.ModePerm); err != nil {
return nil, fmt.Errorf("failed to create data directory for MaxMind database: %w", err)
}
if _, err := os.Stat(service.dbFilePath); os.IsNotExist(err) {
if err := DownloadMaxMindDatabase(service.dbFilePath, downloadURL); err != nil {
return nil, fmt.Errorf("failed to download initial MaxMind database: %w", err)
}
}
if err := service.initialize(); err != nil {
return nil, fmt.Errorf("failed to initialize MaxMind database: %w", err)
}
return service, nil
}
func (s *MaxMindGeoIPService) initialize() error {
s.mu.Lock()
defer s.mu.Unlock()
if s.maxMindDBReader != nil {
_ = s.maxMindDBReader.Close()
s.maxMindDBReader = nil
}
reader, err := maxminddb.Open(s.dbFilePath)
if err != nil {
return fmt.Errorf("error opening MaxMind database at %s: %w", s.dbFilePath, err)
}
s.maxMindDBReader = reader
return nil
}
func (s *MaxMindGeoIPService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
s.mu.RLock()
defer s.mu.RUnlock()
if s.maxMindDBReader == nil {
return nil, fmt.Errorf("MaxMind database is not initialized or failed to open")
}
if ip == nil {
return nil, fmt.Errorf("IP address cannot be nil")
}
var record GeoIpRecord
if err := s.maxMindDBReader.Lookup(ip, &record); err != nil {
return nil, fmt.Errorf("error looking up IP %s in MaxMind database: %w", ip.String(), err)
}
geoInfo := &GeoInfo{
ISOCode: record.Country.ISOCode,
Name: record.Country.Names["en"],
}
if geoInfo.Name == "" && geoInfo.ISOCode != "" {
geoInfo.Name = geoInfo.ISOCode
}
return geoInfo, nil
}
func (s *MaxMindGeoIPService) UpdateDatabase() error {
if err := DownloadMaxMindDatabase(s.dbFilePath, GeoIpUrl); err != nil {
return err
}
return s.initialize()
}
func DownloadMaxMindDatabase(dbFilePath string, downloadURL string) error {
if dbFilePath == "" {
dbFilePath = GeoIpFilePath
}
if downloadURL == "" {
downloadURL = GeoIpUrl
}
resp, err := http.Get(downloadURL)
if err != nil {
return fmt.Errorf("failed to initiate MaxMind database download: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("failed to download MaxMind database: HTTP status %s", resp.Status)
}
if err := os.MkdirAll(filepath.Dir(dbFilePath), os.ModePerm); err != nil {
return fmt.Errorf("failed to create data directory for MaxMind database update: %w", err)
}
tempPath := dbFilePath + ".download"
out, err := os.Create(tempPath)
if err != nil {
return fmt.Errorf("failed to create MaxMind database file at %s: %w", tempPath, err)
}
defer func() {
_ = out.Close()
}()
if _, err = io.Copy(out, resp.Body); err != nil {
return fmt.Errorf("failed to write MaxMind database file: %w", err)
}
if err = out.Close(); err != nil {
return fmt.Errorf("failed to close MaxMind database file: %w", err)
}
if err = os.Rename(tempPath, dbFilePath); err != nil {
return fmt.Errorf("failed to move MaxMind database file into place: %w", err)
}
return nil
}
func (s *MaxMindGeoIPService) Close() error {
s.mu.Lock()
defer s.mu.Unlock()
if s.maxMindDBReader != nil {
err := s.maxMindDBReader.Close()
s.maxMindDBReader = nil
if err != nil {
return fmt.Errorf("error closing MaxMind database: %w", err)
}
}
return nil
}
+152
View File
@@ -0,0 +1,152 @@
package geoip
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"strings"
"time"
"github.com/rain-kl/openflare/openflare-server/utils/geoip/iputil"
)
const defaultOutboundIPLookupTimeout = 5 * time.Second
// OutboundIPStrategy defines a lookup strategy for the current public egress IP.
type OutboundIPStrategy interface {
Name() string
GetOutboundIP(ctx context.Context) (net.IP, error)
}
// OutboundIPAPIAdapter adapts a third-party HTTP API response into an IP value.
type OutboundIPAPIAdapter interface {
Name() string
Endpoint() string
DecodeIP(io.Reader) (net.IP, error)
}
type HTTPOutboundIPStrategy struct {
Client *http.Client
Adapter OutboundIPAPIAdapter
}
func NewHTTPOutboundIPStrategy(adapter OutboundIPAPIAdapter, client *http.Client) *HTTPOutboundIPStrategy {
if client == nil {
client = &http.Client{Timeout: defaultOutboundIPLookupTimeout}
}
return &HTTPOutboundIPStrategy{
Client: client,
Adapter: adapter,
}
}
func (s *HTTPOutboundIPStrategy) Name() string {
if s == nil || s.Adapter == nil {
return "http-outbound-ip"
}
return s.Adapter.Name()
}
func (s *HTTPOutboundIPStrategy) GetOutboundIP(ctx context.Context) (net.IP, error) {
if s == nil || s.Adapter == nil {
return nil, errors.New("outbound IP adapter is nil")
}
if ctx == nil {
ctx = context.Background()
}
client := s.Client
if client == nil {
client = &http.Client{Timeout: defaultOutboundIPLookupTimeout}
}
request, err := http.NewRequestWithContext(ctx, http.MethodGet, s.Adapter.Endpoint(), nil)
if err != nil {
return nil, fmt.Errorf("%s create request failed: %w", s.Name(), err)
}
response, err := client.Do(request)
if err != nil {
return nil, fmt.Errorf("%s request failed: %w", s.Name(), err)
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
return nil, fmt.Errorf("%s returned non-200 status: %d %s", s.Name(), response.StatusCode, response.Status)
}
ip, err := s.Adapter.DecodeIP(response.Body)
if err != nil {
return nil, fmt.Errorf("%s decode response failed: %w", s.Name(), err)
}
if !iputil.IsPublic(ip) {
return nil, fmt.Errorf("%s returned non-public IP: %s", s.Name(), ip.String())
}
return ip, nil
}
type RealIPCCAdapter struct {
URL string
}
type realIPCCResponse struct {
IP string `json:"ip"`
}
func NewRealIPCCOutboundIPStrategy() *HTTPOutboundIPStrategy {
return NewHTTPOutboundIPStrategy(RealIPCCAdapter{}, nil)
}
func (a RealIPCCAdapter) Name() string {
return "realip.cc"
}
func (a RealIPCCAdapter) Endpoint() string {
if strings.TrimSpace(a.URL) != "" {
return strings.TrimSpace(a.URL)
}
return "https://realip.cc"
}
func (a RealIPCCAdapter) DecodeIP(reader io.Reader) (net.IP, error) {
var payload realIPCCResponse
if err := json.NewDecoder(reader).Decode(&payload); err != nil {
return nil, err
}
ip := net.ParseIP(strings.TrimSpace(payload.IP))
if ip == nil {
return nil, fmt.Errorf("invalid IP %q", payload.IP)
}
if ipv4 := ip.To4(); ipv4 != nil {
return ipv4, nil
}
return ip, nil
}
func DefaultOutboundIPStrategies() []OutboundIPStrategy {
return []OutboundIPStrategy{
NewRealIPCCOutboundIPStrategy(),
}
}
func GetOutboundIP(ctx context.Context, strategies ...OutboundIPStrategy) (net.IP, error) {
if len(strategies) == 0 {
strategies = DefaultOutboundIPStrategies()
}
var errs []error
for _, strategy := range strategies {
if strategy == nil {
continue
}
ip, err := strategy.GetOutboundIP(ctx)
if err == nil && ip != nil {
return ip, nil
}
if err != nil {
errs = append(errs, fmt.Errorf("%s: %w", strategy.Name(), err))
}
}
if len(errs) == 0 {
return nil, errors.New("no outbound IP lookup strategy configured")
}
return nil, errors.Join(errs...)
}
@@ -0,0 +1,82 @@
package geoip
import (
"context"
"errors"
"net"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
type fakeOutboundIPStrategy struct {
name string
ip net.IP
err error
}
func (f fakeOutboundIPStrategy) Name() string {
return f.name
}
func (f fakeOutboundIPStrategy) GetOutboundIP(ctx context.Context) (net.IP, error) {
return f.ip, f.err
}
func TestRealIPCCAdapterDecodeIP(t *testing.T) {
ip, err := RealIPCCAdapter{}.DecodeIP(strings.NewReader(`{"ip":"8.8.8.8","country":"United States"}`))
if err != nil {
t.Fatalf("DecodeIP failed: %v", err)
}
if ip.String() != "8.8.8.8" {
t.Fatalf("unexpected IP: %s", ip.String())
}
}
func TestHTTPOutboundIPStrategyUsesAdapter(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
t.Fatalf("unexpected method: %s", r.Method)
}
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"ip":"8.8.4.4"}`))
}))
defer server.Close()
strategy := NewHTTPOutboundIPStrategy(RealIPCCAdapter{URL: server.URL}, server.Client())
ip, err := strategy.GetOutboundIP(context.Background())
if err != nil {
t.Fatalf("GetOutboundIP failed: %v", err)
}
if ip.String() != "8.8.4.4" {
t.Fatalf("unexpected outbound IP: %s", ip.String())
}
}
func TestGetOutboundIPFallsBackToNextStrategy(t *testing.T) {
ip, err := GetOutboundIP(
context.Background(),
fakeOutboundIPStrategy{name: "first", err: errors.New("temporary failure")},
fakeOutboundIPStrategy{name: "second", ip: net.ParseIP("1.1.1.1")},
)
if err != nil {
t.Fatalf("GetOutboundIP failed: %v", err)
}
if ip.String() != "1.1.1.1" {
t.Fatalf("unexpected outbound IP: %s", ip.String())
}
}
func TestHTTPOutboundIPStrategyRejectsPrivateIP(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"ip":"172.17.0.2"}`))
}))
defer server.Close()
strategy := NewHTTPOutboundIPStrategy(RealIPCCAdapter{URL: server.URL}, server.Client())
if _, err := strategy.GetOutboundIP(context.Background()); err == nil {
t.Fatal("expected private IP to be rejected")
}
}