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 }