Files
sub-store/internal/database/repo_extra_test.go
T
rogee de9ac92ede feat: GeoIP-based geo detection + configurable rename template
- Add MaxMind GeoLite2-City mmdb auto-download and IP-based geo lookup
  (covers all 249 countries globally, replacing limited 8-country regex)
- Add RenameOptions with 5 toggleable fields (flag/alias/country/city/index)
  stored as rename_options_json on collections
- Collection rename pipeline uses DetectGeoWithServer(name, server) which
  tries GeoIP lookup first, falls back to name-based regex matching
- Chinese country/city names + emoji flag prefix
- Country-city dedup: skip city when it equals country (e.g. 香港/香港)
- URI fragment encoding: url.QueryEscape -> url.PathEscape (%20 not +)
- Egress cache key excludes name field so renamed nodes still match cache
- Frontend: 5 toggle buttons in collection edit form, default all on
- Migrations: 0004 (source alias + collection rename_enabled),
  0005 (collection rename_options_json)
2026-07-28 14:44:54 +08:00

496 lines
12 KiB
Go

package database
import (
"context"
"path/filepath"
"testing"
"time"
"github.com/jmoiron/sqlx"
_ "modernc.org/sqlite"
"github.com/peterqiu0516/sub-store/internal/model"
)
// --- InitDB ---
func TestInitDB_CreatesDirectoryAndAppliesPragmas(t *testing.T) {
dir := t.TempDir()
dbPath := filepath.Join(dir, "subdir", "test.db")
db, err := InitDB(dbPath)
if err != nil {
t.Fatalf("InitDB failed: %v", err)
}
defer db.Close()
// Verify WAL journal mode
var mode string
if err := db.Get(&mode, "PRAGMA journal_mode"); err != nil {
t.Fatalf("failed to read journal_mode: %v", err)
}
if mode != "wal" {
t.Errorf("expected journal_mode wal, got %s", mode)
}
// Verify foreign_keys
var fk int
if err := db.Get(&fk, "PRAGMA foreign_keys"); err != nil {
t.Fatalf("failed to read foreign_keys: %v", err)
}
if fk != 1 {
t.Errorf("expected foreign_keys 1, got %d", fk)
}
// Verify busy_timeout
var bt int
if err := db.Get(&bt, "PRAGMA busy_timeout"); err != nil {
t.Fatalf("failed to read busy_timeout: %v", err)
}
if bt != 5000 {
t.Errorf("expected busy_timeout 5000, got %d", bt)
}
// Verify synchronous
var sync string
if err := db.Get(&sync, "PRAGMA synchronous"); err != nil {
t.Fatalf("failed to read synchronous: %v", err)
}
if sync != "1" { // NORMAL = 1
t.Errorf("expected synchronous 1 (NORMAL), got %s", sync)
}
// Verify MaxOpenConns
if db.Stats().MaxOpenConnections != 1 {
t.Errorf("expected MaxOpenConnections 1, got %d", db.Stats().MaxOpenConnections)
}
}
func TestInitDB_RelativePath(t *testing.T) {
db, err := InitDB(":memory:")
if err != nil {
// modernc sqlite may not support :memory: via file: DSN; try temp file instead
t.Skipf("InitDB with :memory: not supported: %v", err)
}
defer db.Close()
}
func TestInitDB_ExistingDirectory(t *testing.T) {
dir := t.TempDir()
dbPath := filepath.Join(dir, "test.db")
db, err := InitDB(dbPath)
if err != nil {
t.Fatalf("InitDB failed: %v", err)
}
defer db.Close()
}
// --- StartCacheCleaner ---
func TestStartCacheCleaner_RemovesExpiredEntries(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewCacheRepo(db)
// Insert an expired entry directly (cached_at in the past)
_, err := db.Exec(
`INSERT INTO source_cache (cache_key, content, metadata, cached_at, ttl) VALUES (?, ?, ?, ?, ?)`,
"expired-by-cleaner", "old content", "{}", time.Now().Unix()-100, 1,
)
if err != nil {
t.Fatalf("insert expired entry: %v", err)
}
// Insert a valid entry
repo.SafePut("valid-by-cleaner", "fresh", map[string]any{}, 3600)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
StartCacheCleaner(ctx, db, 100*time.Millisecond)
time.Sleep(350 * time.Millisecond)
cancel()
time.Sleep(50 * time.Millisecond)
// Expired entry should be gone
var count int
db.Get(&count, "SELECT COUNT(*) FROM source_cache WHERE cache_key = ?", "expired-by-cleaner")
if count != 0 {
t.Errorf("expected expired entry to be cleaned, got count=%d", count)
}
// Valid entry should remain
db.Get(&count, "SELECT COUNT(*) FROM source_cache WHERE cache_key = ?", "valid-by-cleaner")
if count != 1 {
t.Errorf("expected valid entry to remain, got count=%d", count)
}
}
// --- SettingsRepo.Replace ---
func TestSettingsReplace(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewSettingsRepo(db)
// Replace with new settings
err := repo.Replace(map[string]any{"appName": "Replaced"})
if err != nil {
t.Fatalf("replace failed: %v", err)
}
settings, _ := repo.Get()
if settings["appName"] != "Replaced" {
t.Errorf("expected appName Replaced, got %v", settings["appName"])
}
// Replace again (should overwrite, not merge)
err = repo.Replace(map[string]any{"theme": "dark"})
if err != nil {
t.Fatalf("replace 2 failed: %v", err)
}
settings, _ = repo.Get()
if _, exists := settings["appName"]; exists {
t.Error("expected appName to be gone after replace (not merge)")
}
if settings["theme"] != "dark" {
t.Errorf("expected theme dark, got %v", settings["theme"])
}
}
// --- CollectionRepo.Delete and Sort ---
func TestCollectionDelete(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewCollectionRepo(db)
repo.Upsert(model.CollectionRecord{
ID: "col-del", Name: "Del", SourceIds: []string{},
TemplateId: "acl4ssr-mihomo", Meta: map[string]any{},
})
err := repo.Delete("col-del")
if err != nil {
t.Fatalf("delete failed: %v", err)
}
got, _ := repo.Get("col-del")
if got != nil {
t.Error("expected nil after delete")
}
}
func TestCollectionSort(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewCollectionRepo(db)
for _, id := range []string{"a", "b", "c"} {
repo.Upsert(model.CollectionRecord{
ID: id, Name: id, SourceIds: []string{},
TemplateId: "acl4ssr-mihomo", Meta: map[string]any{},
})
}
err := repo.Sort([]string{"c", "b", "a"})
if err != nil {
t.Fatalf("sort failed: %v", err)
}
list, _ := repo.List()
if list[0].ID != "c" || list[1].ID != "b" || list[2].ID != "a" {
t.Errorf("sort order wrong: %s, %s, %s", list[0].ID, list[1].ID, list[2].ID)
}
}
// --- TemplateRepo.Delete ---
func TestTemplateDelete(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewTemplateRepo(db)
repo.Upsert(model.TemplateRecord{
ID: "tmpl-del", Name: "Del", Target: "mihomo", Config: map[string]any{},
})
err := repo.Delete("tmpl-del")
if err != nil {
t.Fatalf("delete failed: %v", err)
}
got, _ := repo.Get("tmpl-del")
if got != nil {
t.Error("expected nil after delete")
}
}
// --- RecycleRepo.parseSnapshot edge cases (via List/Get) ---
func TestParseSnapshot_Empty(t *testing.T) {
result := parseSnapshot("")
if len(result) != 0 {
t.Errorf("expected empty map, got %v", result)
}
}
func TestParseSnapshot_InvalidJSON(t *testing.T) {
result := parseSnapshot("not json")
if len(result) != 0 {
t.Errorf("expected empty map for invalid json, got %v", result)
}
}
func TestParseSnapshot_Valid(t *testing.T) {
result := parseSnapshot(`{"key":"value"}`)
if result["key"] != "value" {
t.Errorf("expected key=value, got %v", result)
}
}
func TestRecycleList_Empty(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewRecycleRepo(db)
list, err := repo.List()
if err != nil {
t.Fatalf("list failed: %v", err)
}
if len(list) != 0 {
t.Errorf("expected empty list, got %d", len(list))
}
}
func TestRecycleGet_NotFound(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewRecycleRepo(db)
got, err := repo.Get("nonexistent")
if err != nil {
t.Fatalf("expected nil error, got %v", err)
}
if got != nil {
t.Error("expected nil for not found")
}
}
func TestRecycleArchiveAndDelete_TrimExcess(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewRecycleRepo(db)
// Insert more than MaxRecycleEntries to test trimming
for i := 0; i < 55; i++ {
_, err := repo.ArchiveAndDelete("source", "nonexistent", map[string]any{"i": i}, "SELECT 1", nil)
if err != nil {
t.Fatalf("archive %d failed: %v", i, err)
}
}
list, _ := repo.List()
if len(list) > 50 {
t.Errorf("expected at most 50 entries after trim, got %d", len(list))
}
}
// --- helpers ---
func TestMarshalJSON(t *testing.T) {
if s := marshalJSON(map[string]any{"a": 1}); s != `{"a":1}` {
t.Errorf("expected {\"a\":1}, got %s", s)
}
}
func TestBoolToInt(t *testing.T) {
if boolToInt(true) != 1 {
t.Error("expected 1 for true")
}
if boolToInt(false) != 0 {
t.Error("expected 0 for false")
}
}
func TestJsonUnmarshal(t *testing.T) {
var m map[string]any
jsonUnmarshal(`{"a":1}`, &m)
if m["a"] != float64(1) {
t.Errorf("expected a=1, got %v", m["a"])
}
// Empty string should be a no-op
var m2 map[string]any
jsonUnmarshal("", &m2)
if m2 != nil {
t.Error("expected nil for empty string")
}
}
// --- CacheRepo.Get returns nil on error ---
func TestCacheGet_Miss(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewCacheRepo(db)
entry, err := repo.Get("nonexistent")
if err != nil {
t.Errorf("expected nil error for miss, got %v", err)
}
if entry != nil {
t.Error("expected nil entry for miss")
}
}
func TestCachePut_PanicRecovery(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewCacheRepo(db)
// SafePut should not panic even with nil metadata (normal case)
repo.SafePut("key", "content", nil, 300)
entry, ok := repo.SafeGet("key")
if !ok {
t.Fatal("expected cache hit")
}
if entry.Content != "content" {
t.Errorf("expected content, got %s", entry.Content)
}
if entry.Metadata == nil {
t.Error("expected non-nil metadata")
}
}
func TestCacheGet_Expired(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewCacheRepo(db)
repo.Put("expired", "old", map[string]any{}, 1)
time.Sleep(2 * time.Second)
entry, err := repo.Get("expired")
if err != nil {
t.Errorf("expected nil error, got %v", err)
}
if entry != nil {
t.Error("expected nil for expired")
}
}
// --- SourceRepo additional coverage ---
func TestSourceUpsert_AutoId(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewSourceRepo(db)
rec, err := repo.Upsert(model.SourceRecord{
Name: "My Source", Type: "local", Content: "ss://pass@host:port#name",
Enabled: true, Meta: map[string]any{},
})
if err != nil {
t.Fatalf("upsert failed: %v", err)
}
if rec.ID != "my-source" {
t.Errorf("expected id my-source, got %s", rec.ID)
}
}
func TestSourceUpsert_DefaultTypeRemote(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewSourceRepo(db)
repo.Upsert(model.SourceRecord{
ID: "test", Name: "Test", Type: "", URL: "https://example.com/sub",
Enabled: true, Meta: map[string]any{},
})
// Read back — sourceFromRow normalizes empty type to "remote"
got, _ := repo.Get("test")
if got.Type != "remote" {
t.Errorf("expected type remote, got %s", got.Type)
}
}
func TestSourceGet_NotFound(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewSourceRepo(db)
got, err := repo.Get("nonexistent")
if err != nil {
t.Fatalf("expected nil error, got %v", err)
}
if got != nil {
t.Error("expected nil for not found")
}
}
// --- TemplateRepo additional coverage ---
func TestTemplateUpsert_DefaultTargetMihomo(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewTemplateRepo(db)
rec, _ := repo.Upsert(model.TemplateRecord{
ID: "test", Name: "Test", Target: "", Config: nil,
})
if rec.Target != "mihomo" {
t.Errorf("expected target mihomo, got %s", rec.Target)
}
if rec.Config == nil {
t.Error("expected non-nil config")
}
}
func TestTemplateGet_NotFound(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewTemplateRepo(db)
got, err := repo.Get("nonexistent")
if err != nil {
t.Fatalf("expected nil error, got %v", err)
}
if got != nil {
t.Error("expected nil for not found")
}
}
func TestTemplateFromRow_EmptyTarget(t *testing.T) {
rec := templateFromRow(templateRow{ID: "x", Target: "", ConfigJSON: "{}"})
if rec.Target != "mihomo" {
t.Errorf("expected mihomo, got %s", rec.Target)
}
}
// --- CollectionRepo additional coverage ---
func TestCollectionUpsert_Defaults(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewCollectionRepo(db)
rec, _ := repo.Upsert(model.CollectionRecord{
ID: "test", Name: "Test",
})
if rec.TemplateId != "acl4ssr-mihomo" {
t.Errorf("expected default template, got %s", rec.TemplateId)
}
if rec.SourceIds == nil {
t.Error("expected non-nil SourceIds")
}
if rec.Meta == nil {
t.Error("expected non-nil Meta")
}
}
func TestCollectionGet_NotFound(t *testing.T) {
db := testDB(t)
defer db.Close()
repo := NewCollectionRepo(db)
got, err := repo.Get("nonexistent")
if err != nil {
t.Fatalf("expected nil error, got %v", err)
}
if got != nil {
t.Error("expected nil for not found")
}
}
// --- RunMigrations is already tested via testDB; test error case ---
func TestRunMigrations_OnFreshDB(t *testing.T) {
dir := t.TempDir()
dbPath := filepath.Join(dir, "fresh.db")
db, err := sqlx.Open("sqlite", dbPath)
if err != nil {
t.Fatalf("open failed: %v", err)
}
defer db.Close()
if err := RunMigrations(db); err != nil {
t.Fatalf("migrations failed: %v", err)
}
// Tables should exist
var count int
db.Get(&count, "SELECT COUNT(*) FROM sources")
if count != 0 {
t.Errorf("expected 0 sources, got %d", count)
}
}