Files
OpenGFW/ruleset/builtins/geo/geo_matcher.go
T

178 lines
4.3 KiB
Go
Raw Normal View History

package geo
import (
"net"
"sort"
"strings"
"sync"
)
type GeoMatcher struct {
geoLoader GeoLoader
geoSiteMatcher map[string]hostMatcher
siteMatcherLock sync.Mutex
geoIpMatcher map[string]hostMatcher
ipMatcherLock sync.Mutex
}
func NewGeoMatcher(geoSiteFilename, geoIpFilename string) *GeoMatcher {
return &GeoMatcher{
geoLoader: NewDefaultGeoLoader(geoSiteFilename, geoIpFilename),
geoSiteMatcher: make(map[string]hostMatcher),
geoIpMatcher: make(map[string]hostMatcher),
}
}
func (g *GeoMatcher) MatchGeoIp(ip, condition string) bool {
g.ipMatcherLock.Lock()
defer g.ipMatcherLock.Unlock()
matcher, ok := g.geoIpMatcher[condition]
if !ok {
// GeoIP matcher
condition = strings.ToLower(condition)
country := condition
if len(country) == 0 {
return false
}
gMap, err := g.geoLoader.LoadGeoIP()
if err != nil {
return false
}
list, ok := gMap[country]
if !ok || list == nil {
return false
}
matcher, err = newGeoIPMatcher(list)
if err != nil {
return false
}
g.geoIpMatcher[condition] = matcher
}
parseIp := net.ParseIP(ip)
if parseIp == nil {
return false
}
ipv4 := parseIp.To4()
if ipv4 != nil {
return matcher.Match(HostInfo{IPv4: ipv4})
}
ipv6 := parseIp.To16()
if ipv6 != nil {
return matcher.Match(HostInfo{IPv6: ipv6})
}
return false
}
func (g *GeoMatcher) MatchGeoSite(site, condition string) bool {
g.siteMatcherLock.Lock()
defer g.siteMatcherLock.Unlock()
matcher, ok := g.geoSiteMatcher[condition]
if !ok {
// MatchGeoSite matcher
condition = strings.ToLower(condition)
name, attrs := parseGeoSiteName(condition)
if len(name) == 0 {
return false
}
gMap, err := g.geoLoader.LoadGeoSite()
if err != nil {
return false
}
list, ok := gMap[name]
if !ok || list == nil {
return false
}
matcher, err = newGeositeMatcher(list, attrs)
if err != nil {
return false
}
g.geoSiteMatcher[condition] = matcher
}
return matcher.Match(HostInfo{Name: site})
}
func (g *GeoMatcher) LoadGeoSite() error {
_, err := g.geoLoader.LoadGeoSite()
return err
}
// GeoIPEntry describes one entry of a GeoIP database.
type GeoIPEntry struct {
// Code is the lowercase key to pass to geoip(), usually a country code.
Code string
// CIDRs is the number of networks the entry covers.
CIDRs int
}
// GeoSiteEntry describes one entry of a GeoSite database.
type GeoSiteEntry struct {
// Code is the lowercase key to pass to geosite().
Code string
// Attributes are the suffixes usable as `code@attribute`.
Attributes []string
// Domains is the number of domain rules the entry covers.
Domains int
}
// ListGeoIP returns every entry of the GeoIP database, sorted by code.
// It is meant for UIs that let the user pick a country instead of typing one.
func (g *GeoMatcher) ListGeoIP() ([]GeoIPEntry, error) {
gMap, err := g.geoLoader.LoadGeoIP()
if err != nil {
return nil, err
}
entries := make([]GeoIPEntry, 0, len(gMap))
for code, list := range gMap {
entries = append(entries, GeoIPEntry{Code: code, CIDRs: len(list.GetCidr())})
}
sort.Slice(entries, func(i, j int) bool { return entries[i].Code < entries[j].Code })
return entries, nil
}
// ListGeoSite returns every entry of the GeoSite database, sorted by code,
// including the attributes each entry supports.
func (g *GeoMatcher) ListGeoSite() ([]GeoSiteEntry, error) {
gMap, err := g.geoLoader.LoadGeoSite()
if err != nil {
return nil, err
}
entries := make([]GeoSiteEntry, 0, len(gMap))
for code, list := range gMap {
attrSet := make(map[string]struct{})
for _, domain := range list.GetDomain() {
for _, attr := range domain.GetAttribute() {
attrSet[strings.ToLower(attr.GetKey())] = struct{}{}
}
}
attrs := make([]string, 0, len(attrSet))
for attr := range attrSet {
attrs = append(attrs, attr)
}
sort.Strings(attrs)
entries = append(entries, GeoSiteEntry{
Code: code,
Attributes: attrs,
Domains: len(list.GetDomain()),
})
}
sort.Slice(entries, func(i, j int) bool { return entries[i].Code < entries[j].Code })
return entries, nil
}
func (g *GeoMatcher) LoadGeoIP() error {
_, err := g.geoLoader.LoadGeoIP()
return err
}
func parseGeoSiteName(s string) (string, []string) {
parts := strings.Split(s, "@")
base := strings.TrimSpace(parts[0])
attrs := parts[1:]
for i := range attrs {
attrs[i] = strings.TrimSpace(attrs[i])
}
return base, attrs
}