geodata_test.go 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558
  1. package geodata
  2. import (
  3. "encoding/json"
  4. "errors"
  5. "net/netip"
  6. "os"
  7. "path/filepath"
  8. "strconv"
  9. "strings"
  10. "sync"
  11. "testing"
  12. "time"
  13. xraygeodata "github.com/xtls/xray-core/common/geodata"
  14. "google.golang.org/protobuf/proto"
  15. )
  16. func writeSiteDB(t *testing.T, dir, name string, sites ...*xraygeodata.GeoSite) string {
  17. t.Helper()
  18. data, err := proto.Marshal(&xraygeodata.GeoSiteList{Entry: sites})
  19. if err != nil {
  20. t.Fatalf("marshal geosite list: %v", err)
  21. }
  22. return writeFile(t, dir, name, data)
  23. }
  24. func writeIPDB(t *testing.T, dir, name string, geoips ...*xraygeodata.GeoIP) string {
  25. t.Helper()
  26. data, err := proto.Marshal(&xraygeodata.GeoIPList{Entry: geoips})
  27. if err != nil {
  28. t.Fatalf("marshal geoip list: %v", err)
  29. }
  30. return writeFile(t, dir, name, data)
  31. }
  32. func writeFile(t *testing.T, dir, name string, data []byte) string {
  33. t.Helper()
  34. path := filepath.Join(dir, name)
  35. if err := os.WriteFile(path, data, 0o644); err != nil {
  36. t.Fatalf("write %s: %v", name, err)
  37. }
  38. return path
  39. }
  40. func site(code string, domains ...*xraygeodata.Domain) *xraygeodata.GeoSite {
  41. return &xraygeodata.GeoSite{Code: code, Domain: domains}
  42. }
  43. func domain(domainType xraygeodata.Domain_Type, value string, attributes ...string) *xraygeodata.Domain {
  44. d := &xraygeodata.Domain{Type: domainType, Value: value}
  45. for _, attribute := range attributes {
  46. d.Attribute = append(d.Attribute, &xraygeodata.Domain_Attribute{
  47. Key: attribute,
  48. TypedValue: &xraygeodata.Domain_Attribute_BoolValue{BoolValue: true},
  49. })
  50. }
  51. return d
  52. }
  53. func geoip(code string, prefixes ...string) *xraygeodata.GeoIP {
  54. entry := &xraygeodata.GeoIP{Code: code}
  55. for _, raw := range prefixes {
  56. prefix := netip.MustParsePrefix(raw)
  57. entry.Cidr = append(entry.Cidr, &xraygeodata.CIDR{
  58. Ip: prefix.Addr().AsSlice(),
  59. Prefix: uint32(prefix.Bits()),
  60. })
  61. }
  62. return entry
  63. }
  64. func sampleSiteDB(t *testing.T, dir string) {
  65. t.Helper()
  66. writeSiteDB(t, dir, "geosite.dat",
  67. site("google",
  68. domain(xraygeodata.Domain_Domain, "google.com"),
  69. domain(xraygeodata.Domain_Full, "ads.google.com", "ads"),
  70. domain(xraygeodata.Domain_Substr, "googlevideo", "cn"),
  71. domain(xraygeodata.Domain_Regex, `^g.*\.cn$`),
  72. ),
  73. site("CN",
  74. domain(xraygeodata.Domain_Domain, "baidu.com"),
  75. domain(xraygeodata.Domain_Domain, "qq.com"),
  76. ),
  77. )
  78. }
  79. func TestListFilesReportsKindAndCategories(t *testing.T) {
  80. dir := t.TempDir()
  81. sampleSiteDB(t, dir)
  82. writeIPDB(t, dir, "geoip.dat", geoip("cn", "1.0.1.0/24"), geoip("private", "10.0.0.0/8", "fc00::/7"))
  83. files, err := NewStore(dir).ListFiles()
  84. if err != nil {
  85. t.Fatalf("ListFiles() error = %v", err)
  86. }
  87. if len(files) != 2 {
  88. t.Fatalf("ListFiles() returned %d files, want 2", len(files))
  89. }
  90. byName := make(map[string]GeoFile, len(files))
  91. for _, file := range files {
  92. byName[file.Name] = file
  93. }
  94. geosite := byName["geosite.dat"]
  95. if geosite.Kind != KindSite {
  96. t.Errorf("geosite.dat kind = %q, want %q", geosite.Kind, KindSite)
  97. }
  98. if geosite.Categories != 2 {
  99. t.Errorf("geosite.dat categories = %d, want 2", geosite.Categories)
  100. }
  101. if geosite.Error != "" {
  102. t.Errorf("geosite.dat error = %q, want empty", geosite.Error)
  103. }
  104. geoipFile := byName["geoip.dat"]
  105. if geoipFile.Kind != KindIP {
  106. t.Errorf("geoip.dat kind = %q, want %q", geoipFile.Kind, KindIP)
  107. }
  108. if geoipFile.Categories != 2 {
  109. t.Errorf("geoip.dat categories = %d, want 2", geoipFile.Categories)
  110. }
  111. }
  112. func TestKindDetectedFromContentsNotName(t *testing.T) {
  113. dir := t.TempDir()
  114. writeSiteDB(t, dir, "my_ip_rules.dat", site("corp", domain(xraygeodata.Domain_Domain, "intranet.corp.local")))
  115. writeIPDB(t, dir, "custom_sites.dat", geoip("office", "192.168.7.0/24"))
  116. store := NewStore(dir)
  117. sitePage, err := store.Categories("my_ip_rules.dat", "", 0, 10)
  118. if err != nil {
  119. t.Fatalf("Categories(my_ip_rules.dat) error = %v", err)
  120. }
  121. if sitePage.Total != 1 || sitePage.Items[0].Code != "corp" {
  122. t.Fatalf("Categories(my_ip_rules.dat) = %+v, want single category corp", sitePage)
  123. }
  124. entries, err := store.Entries("custom_sites.dat", "office", "", 0, 10)
  125. if err != nil {
  126. t.Fatalf("Entries(custom_sites.dat) error = %v", err)
  127. }
  128. if len(entries.Items) != 1 {
  129. t.Fatalf("Entries(custom_sites.dat) returned %d items, want 1", len(entries.Items))
  130. }
  131. if got := entries.Items[0]; got.Kind != "cidr" || got.Value != "192.168.7.0/24" {
  132. t.Errorf("entry = %+v, want cidr 192.168.7.0/24", got)
  133. }
  134. }
  135. func TestEntriesMapDomainTypesAndAttributes(t *testing.T) {
  136. dir := t.TempDir()
  137. sampleSiteDB(t, dir)
  138. store := NewStore(dir)
  139. page, err := store.Entries("geosite.dat", "google", "", 0, 10)
  140. if err != nil {
  141. t.Fatalf("Entries() error = %v", err)
  142. }
  143. want := []GeoEntry{
  144. {Kind: "domain", Value: "google.com"},
  145. {Kind: "full", Value: "ads.google.com"},
  146. {Kind: "keyword", Value: "googlevideo"},
  147. {Kind: "regexp", Value: `^g.*\.cn$`},
  148. }
  149. if page.Total != len(want) {
  150. t.Fatalf("Entries() total = %d, want %d", page.Total, len(want))
  151. }
  152. for i, entry := range want {
  153. if page.Items[i] != entry {
  154. t.Errorf("entry %d = %+v, want %+v", i, page.Items[i], entry)
  155. }
  156. }
  157. category, err := store.Lookup("geosite.dat", "google")
  158. if err != nil {
  159. t.Fatalf("Lookup() error = %v", err)
  160. }
  161. if len(category.Attributes) != 2 || category.Attributes[0] != "ads" || category.Attributes[1] != "cn" {
  162. t.Errorf("attributes = %v, want [ads cn]", category.Attributes)
  163. }
  164. }
  165. func TestCategoriesWithoutAttributesMarshalAsEmptyArray(t *testing.T) {
  166. dir := t.TempDir()
  167. sampleSiteDB(t, dir)
  168. page, err := NewStore(dir).Categories("geosite.dat", "cn", 0, 10)
  169. if err != nil {
  170. t.Fatalf("Categories() error = %v", err)
  171. }
  172. if page.Items[0].Attributes == nil {
  173. t.Fatal("attributes are nil, want an empty slice so the JSON stays an array")
  174. }
  175. encoded, err := json.Marshal(page.Items[0])
  176. if err != nil {
  177. t.Fatalf("marshal category: %v", err)
  178. }
  179. if !strings.Contains(string(encoded), `"attributes":[]`) {
  180. t.Errorf("encoded category = %s, want an empty attributes array", encoded)
  181. }
  182. }
  183. func TestCategoryCodesAreLowercasedAndSorted(t *testing.T) {
  184. dir := t.TempDir()
  185. sampleSiteDB(t, dir)
  186. page, err := NewStore(dir).Categories("geosite.dat", "", 0, 10)
  187. if err != nil {
  188. t.Fatalf("Categories() error = %v", err)
  189. }
  190. if page.Items[0].Code != "cn" || page.Items[1].Code != "google" {
  191. t.Errorf("codes = %q, %q; want cn, google", page.Items[0].Code, page.Items[1].Code)
  192. }
  193. }
  194. func TestSearchFilters(t *testing.T) {
  195. dir := t.TempDir()
  196. sampleSiteDB(t, dir)
  197. store := NewStore(dir)
  198. categories, err := store.Categories("geosite.dat", "OOG", 0, 10)
  199. if err != nil {
  200. t.Fatalf("Categories() error = %v", err)
  201. }
  202. if categories.Total != 1 || categories.Items[0].Code != "google" {
  203. t.Errorf("Categories(OOG) = %+v, want only google", categories)
  204. }
  205. entries, err := store.Entries("geosite.dat", "google", "ADS.", 0, 10)
  206. if err != nil {
  207. t.Fatalf("Entries() error = %v", err)
  208. }
  209. if entries.Total != 1 || entries.Items[0].Value != "ads.google.com" {
  210. t.Errorf("Entries(ADS.) = %+v, want only ads.google.com", entries)
  211. }
  212. }
  213. func TestPagination(t *testing.T) {
  214. dir := t.TempDir()
  215. domains := make([]*xraygeodata.Domain, 0, 250)
  216. for i := range 250 {
  217. domains = append(domains, domain(xraygeodata.Domain_Domain, "host"+strconv.Itoa(i)+".example.com"))
  218. }
  219. writeSiteDB(t, dir, "geosite.dat", site("bulk", domains...))
  220. store := NewStore(dir)
  221. tests := []struct {
  222. name string
  223. offset int
  224. limit int
  225. wantCount int
  226. wantFirst string
  227. }{
  228. {name: "first page", offset: 0, limit: 10, wantCount: 10, wantFirst: "host0.example.com"},
  229. {name: "middle page", offset: 20, limit: 5, wantCount: 5, wantFirst: "host20.example.com"},
  230. {name: "negative offset clamps to start", offset: -5, limit: 3, wantCount: 3, wantFirst: "host0.example.com"},
  231. {name: "tail shorter than limit", offset: 245, limit: 50, wantCount: 5, wantFirst: "host245.example.com"},
  232. {name: "offset past end", offset: 900, limit: 10, wantCount: 0},
  233. {name: "limit above cap", offset: 0, limit: 5000, wantCount: 250, wantFirst: "host0.example.com"},
  234. {name: "zero limit uses cap", offset: 0, limit: 0, wantCount: 250, wantFirst: "host0.example.com"},
  235. }
  236. for _, tt := range tests {
  237. t.Run(tt.name, func(t *testing.T) {
  238. page, err := store.Entries("geosite.dat", "bulk", "", tt.offset, tt.limit)
  239. if err != nil {
  240. t.Fatalf("Entries() error = %v", err)
  241. }
  242. if page.Total != 250 {
  243. t.Errorf("total = %d, want 250", page.Total)
  244. }
  245. if len(page.Items) != tt.wantCount {
  246. t.Fatalf("items = %d, want %d", len(page.Items), tt.wantCount)
  247. }
  248. if tt.wantFirst != "" && page.Items[0].Value != tt.wantFirst {
  249. t.Errorf("first item = %q, want %q", page.Items[0].Value, tt.wantFirst)
  250. }
  251. })
  252. }
  253. }
  254. func TestCategoriesReturnEverythingWithoutLimit(t *testing.T) {
  255. dir := t.TempDir()
  256. sites := make([]*xraygeodata.GeoSite, 0, MaxPageSize+20)
  257. for i := range MaxPageSize + 20 {
  258. sites = append(sites, site("cat"+strconv.Itoa(i), domain(xraygeodata.Domain_Domain, "example.com")))
  259. }
  260. writeSiteDB(t, dir, "geosite.dat", sites...)
  261. store := NewStore(dir)
  262. all, err := store.Categories("geosite.dat", "", 0, 0)
  263. if err != nil {
  264. t.Fatalf("Categories() error = %v", err)
  265. }
  266. if len(all.Items) != MaxPageSize+20 {
  267. t.Errorf("items without a limit = %d, want %d", len(all.Items), MaxPageSize+20)
  268. }
  269. capped, err := store.Categories("geosite.dat", "", 0, 10)
  270. if err != nil {
  271. t.Fatalf("Categories() error = %v", err)
  272. }
  273. if len(capped.Items) != 10 || capped.Total != MaxPageSize+20 {
  274. t.Errorf("explicit limit gave %d items with total %d, want 10 and %d", len(capped.Items), capped.Total, MaxPageSize+20)
  275. }
  276. entries, err := store.Entries("geosite.dat", "cat0", "", 0, 0)
  277. if err != nil {
  278. t.Fatalf("Entries() error = %v", err)
  279. }
  280. if len(entries.Items) != 1 {
  281. t.Errorf("entries = %d, want 1", len(entries.Items))
  282. }
  283. }
  284. func TestErrors(t *testing.T) {
  285. dir := t.TempDir()
  286. sampleSiteDB(t, dir)
  287. writeFile(t, dir, "broken.dat", []byte("this is not a protobuf message at all"))
  288. store := NewStore(dir)
  289. tests := []struct {
  290. name string
  291. call func() error
  292. want error
  293. }{
  294. {
  295. name: "unknown category",
  296. call: func() error { _, err := store.Entries("geosite.dat", "nope", "", 0, 10); return err },
  297. want: ErrUnknownCategory,
  298. },
  299. {
  300. name: "lookup of unknown category",
  301. call: func() error { _, err := store.Lookup("geosite.dat", "nope"); return err },
  302. want: ErrUnknownCategory,
  303. },
  304. {
  305. name: "path traversal",
  306. call: func() error { _, err := store.Categories("../geosite.dat", "", 0, 10); return err },
  307. want: ErrInvalidName,
  308. },
  309. {
  310. name: "non dat extension",
  311. call: func() error { _, err := store.Categories("x-ui.db", "", 0, 10); return err },
  312. want: ErrInvalidName,
  313. },
  314. {
  315. name: "empty name",
  316. call: func() error { _, err := store.Categories("", "", 0, 10); return err },
  317. want: ErrInvalidName,
  318. },
  319. {
  320. name: "unparsable file",
  321. call: func() error { _, err := store.Categories("broken.dat", "", 0, 10); return err },
  322. want: ErrUnrecognized,
  323. },
  324. }
  325. for _, tt := range tests {
  326. t.Run(tt.name, func(t *testing.T) {
  327. if err := tt.call(); !errors.Is(err, tt.want) {
  328. t.Errorf("error = %v, want %v", err, tt.want)
  329. }
  330. })
  331. }
  332. }
  333. func TestBrokenFileIsListedWithReason(t *testing.T) {
  334. dir := t.TempDir()
  335. writeFile(t, dir, "broken.dat", []byte("not a database"))
  336. files, err := NewStore(dir).ListFiles()
  337. if err != nil {
  338. t.Fatalf("ListFiles() error = %v", err)
  339. }
  340. if len(files) != 1 {
  341. t.Fatalf("ListFiles() returned %d files, want 1", len(files))
  342. }
  343. if !strings.HasPrefix(files[0].Error, ErrUnrecognized.Error()) {
  344. t.Errorf("error = %q, want it to start with %q", files[0].Error, ErrUnrecognized.Error())
  345. }
  346. if files[0].Kind != "" {
  347. t.Errorf("kind = %q, want empty", files[0].Kind)
  348. }
  349. }
  350. func TestFileAboveSizeLimitIsRejected(t *testing.T) {
  351. dir := t.TempDir()
  352. path := writeFile(t, dir, "huge.dat", []byte("x"))
  353. if err := os.Truncate(path, MaxFileSize+1); err != nil {
  354. t.Fatalf("truncate: %v", err)
  355. }
  356. store := NewStore(dir)
  357. if _, err := store.Categories("huge.dat", "", 0, 10); !errors.Is(err, ErrFileTooLarge) {
  358. t.Errorf("error = %v, want %v", err, ErrFileTooLarge)
  359. }
  360. files, err := store.ListFiles()
  361. if err != nil {
  362. t.Fatalf("ListFiles() error = %v", err)
  363. }
  364. if len(files) != 1 || files[0].Error != ErrFileTooLarge.Error() {
  365. t.Errorf("ListFiles() = %+v, want the file listed with a too-large error", files)
  366. }
  367. }
  368. func TestIndexCacheInvalidatedWhenFileChanges(t *testing.T) {
  369. dir := t.TempDir()
  370. sampleSiteDB(t, dir)
  371. store := NewStore(dir)
  372. before, err := store.Categories("geosite.dat", "", 0, 10)
  373. if err != nil {
  374. t.Fatalf("Categories() error = %v", err)
  375. }
  376. if before.Total != 2 {
  377. t.Fatalf("total before rewrite = %d, want 2", before.Total)
  378. }
  379. path := writeSiteDB(t, dir, "geosite.dat",
  380. site("google", domain(xraygeodata.Domain_Domain, "google.com")),
  381. site("cn", domain(xraygeodata.Domain_Domain, "baidu.com")),
  382. site("telegram", domain(xraygeodata.Domain_Domain, "t.me")),
  383. )
  384. touch(t, path, time.Now().Add(time.Second))
  385. after, err := store.Categories("geosite.dat", "", 0, 10)
  386. if err != nil {
  387. t.Fatalf("Categories() after rewrite error = %v", err)
  388. }
  389. if after.Total != 3 {
  390. t.Errorf("total after rewrite = %d, want 3", after.Total)
  391. }
  392. if len(store.indexes) != 1 {
  393. t.Errorf("cached indexes = %d, want 1 after the stale entry is dropped", len(store.indexes))
  394. }
  395. }
  396. func touch(t *testing.T, path string, when time.Time) {
  397. t.Helper()
  398. if err := os.Chtimes(path, when, when); err != nil {
  399. t.Fatalf("chtimes %s: %v", path, err)
  400. }
  401. }
  402. func TestDefaultRouteCIDRSurvives(t *testing.T) {
  403. dir := t.TempDir()
  404. writeIPDB(t, dir, "geoip.dat", geoip("any", "0.0.0.0/0", "::/0"), geoip("cn", "1.0.1.0/24"))
  405. page, err := NewStore(dir).Entries("geoip.dat", "any", "", 0, 10)
  406. if err != nil {
  407. t.Fatalf("Entries() error = %v", err)
  408. }
  409. if page.Total != 2 {
  410. t.Fatalf("total = %d, want 2 — a zero prefix is omitted by proto3 and must not be dropped", page.Total)
  411. }
  412. if page.Items[0].Value != "0.0.0.0/0" || page.Items[1].Value != "::/0" {
  413. t.Errorf("items = %+v, want the two default routes", page.Items)
  414. }
  415. }
  416. func TestBrokenFileIsParsedOnlyOnce(t *testing.T) {
  417. dir := t.TempDir()
  418. writeFile(t, dir, "broken.dat", []byte("not a database"))
  419. store := NewStore(dir)
  420. for range 3 {
  421. if _, err := store.Categories("broken.dat", "", 0, 10); !errors.Is(err, ErrUnrecognized) {
  422. t.Fatalf("error = %v, want %v", err, ErrUnrecognized)
  423. }
  424. }
  425. if len(store.indexes) != 1 {
  426. t.Errorf("cached indexes = %d, want the failure cached once", len(store.indexes))
  427. }
  428. }
  429. func TestConcurrentReadsAreConsistent(t *testing.T) {
  430. dir := t.TempDir()
  431. sampleSiteDB(t, dir)
  432. writeIPDB(t, dir, "geoip.dat", geoip("private", "10.0.0.0/8"))
  433. store := NewStore(dir)
  434. var wg sync.WaitGroup
  435. for i := range 24 {
  436. wg.Add(1)
  437. go func(worker int) {
  438. defer wg.Done()
  439. switch worker % 3 {
  440. case 0:
  441. page, err := store.Categories("geosite.dat", "", 0, 0)
  442. if err != nil || page.Total != 2 {
  443. t.Errorf("Categories() = %+v, err = %v; want 2 categories", page, err)
  444. }
  445. case 1:
  446. page, err := store.Entries("geosite.dat", "google", "", 0, 10)
  447. if err != nil || page.Total != 4 {
  448. t.Errorf("Entries() = %+v, err = %v; want 4 entries", page, err)
  449. }
  450. default:
  451. files, err := store.ListFiles()
  452. if err != nil || len(files) != 2 {
  453. t.Errorf("ListFiles() = %d files, err = %v; want 2 files", len(files), err)
  454. }
  455. }
  456. }(i)
  457. }
  458. wg.Wait()
  459. }
  460. func TestLookupDoesNotForgiveStraySpaces(t *testing.T) {
  461. dir := t.TempDir()
  462. sampleSiteDB(t, dir)
  463. store := NewStore(dir)
  464. if _, err := store.Lookup("geosite.dat", "google"); err != nil {
  465. t.Fatalf("Lookup(google) error = %v", err)
  466. }
  467. for _, code := range []string{" google", "google ", "goo gle"} {
  468. if _, err := store.Lookup("geosite.dat", code); !errors.Is(err, ErrUnknownCategory) {
  469. t.Errorf("Lookup(%q) error = %v, want %v — the core does not trim either", code, err, ErrUnknownCategory)
  470. }
  471. }
  472. }
  473. func TestSymlinkOutOfTheAssetFolderIsRefused(t *testing.T) {
  474. outside := t.TempDir()
  475. secret := filepath.Join(outside, "secret.dat")
  476. if err := os.WriteFile(secret, []byte("not yours"), 0o644); err != nil {
  477. t.Fatalf("write secret: %v", err)
  478. }
  479. dir := t.TempDir()
  480. sampleSiteDB(t, dir)
  481. if err := os.Symlink(secret, filepath.Join(dir, "escape.dat")); err != nil {
  482. t.Skipf("symlinks unavailable: %v", err)
  483. }
  484. store := NewStore(dir)
  485. if _, err := store.Categories("escape.dat", "", 0, 10); err == nil {
  486. t.Error("Categories() read through a symlink pointing outside the asset folder")
  487. }
  488. if _, err := store.Entries("escape.dat", "google", "", 0, 10); err == nil {
  489. t.Error("Entries() read through a symlink pointing outside the asset folder")
  490. }
  491. files, err := store.ListFiles()
  492. if err != nil {
  493. t.Fatalf("ListFiles() error = %v", err)
  494. }
  495. for _, file := range files {
  496. if file.Name == "escape.dat" && file.Error == "" {
  497. t.Error("ListFiles() reported an escaping symlink as a usable database")
  498. }
  499. }
  500. }