Browse Source

Fix clients group filter and search for non-ASCII capitals (#6685)

* Match non-ASCII capitals in the clients group filter and search

SQLite LOWER() only folds ASCII, so a group or search term with a Cyrillic,
Persian or other non-ASCII capital never matched after being lower-cased in
Go. Also compare the value as typed.

* Match lower, upper and title-case spellings for non-ASCII search and group filter
kaveh 13 hours ago
parent
commit
7802167443
2 changed files with 93 additions and 11 deletions
  1. 60 11
      internal/web/service/client_paging.go
  2. 33 0
      internal/web/service/client_paging_test.go

+ 60 - 11
internal/web/service/client_paging.go

@@ -1,6 +1,7 @@
 package service
 
 import (
+	"slices"
 	"sort"
 	"strconv"
 	"strings"
@@ -113,13 +114,45 @@ const (
 	sqlClientEnabled = "COALESCE(c.enable, FALSE)"
 )
 
-const clientSearchCond = `(LOWER(c.email) LIKE ? ESCAPE '\'
-	OR LOWER(COALESCE(c.sub_id, '')) LIKE ? ESCAPE '\'
-	OR LOWER(COALESCE(c.comment, '')) LIKE ? ESCAPE '\'
-	OR LOWER(COALESCE(c.uuid, '')) LIKE ? ESCAPE '\'
-	OR LOWER(COALESCE(c.password, '')) LIKE ? ESCAPE '\'
-	OR LOWER(COALESCE(c.auth, '')) LIKE ? ESCAPE '\'
-	OR (COALESCE(c.tg_id, 0) <> 0 AND CAST(c.tg_id AS TEXT) LIKE ? ESCAPE '\'))`
+// clientSearchCols are the text columns the search box matches.
+var clientSearchCols = []string{"c.email", "COALESCE(c.sub_id, '')", "COALESCE(c.comment, '')",
+	"COALESCE(c.uuid, '')", "COALESCE(c.password, '')", "COALESCE(c.auth, '')"}
+
+// caseVariants returns s lower-cased, as typed, upper-cased and title-cased.
+// SQLite's LOWER() and LIKE fold ASCII only, so non-ASCII text is matched
+// against these spellings instead of relying on the database to fold it.
+func caseVariants(s string) []string {
+	title := s
+	if r := []rune(strings.ToLower(s)); len(r) > 0 {
+		title = strings.ToUpper(string(r[:1])) + string(r[1:])
+	}
+	out := make([]string, 0, 4)
+	for _, v := range []string{strings.ToLower(s), s, strings.ToUpper(s), title} {
+		if !slices.Contains(out, v) {
+			out = append(out, v)
+		}
+	}
+	return out
+}
+
+// clientSearchCond builds the search predicate for the given needle variants
+// (lowered first) and returns its arguments.
+func clientSearchCond(variants []string) (string, []any) {
+	var parts []string
+	var args []any
+	like := func(v string) string { return "%" + escapeLikeLiteral(v) + "%" }
+	for _, col := range clientSearchCols {
+		parts = append(parts, "LOWER("+col+") LIKE ? ESCAPE '\\'")
+		args = append(args, like(variants[0]))
+		for _, v := range variants {
+			parts = append(parts, col+" LIKE ? ESCAPE '\\'")
+			args = append(args, like(v))
+		}
+	}
+	parts = append(parts, "(COALESCE(c.tg_id, 0) <> 0 AND CAST(c.tg_id AS TEXT) LIKE ? ESCAPE '\\')")
+	args = append(args, like(variants[0]))
+	return "(" + strings.Join(parts, " OR ") + ")", args
+}
 
 // clientQuery builds the statements behind the clients page: a clients row
 // joined to its traffic counters, plus the expressions every bucket predicate
@@ -213,9 +246,9 @@ func (q clientQuery) applyParams(tx *gorm.DB, params ClientPageParams, onlines [
 		tx = tx.Where(cond, args...)
 	}
 
-	if needle := strings.ToLower(strings.TrimSpace(params.Search)); needle != "" {
-		pattern := "%" + escapeLikeLiteral(needle) + "%"
-		where(clientSearchCond, pattern, pattern, pattern, pattern, pattern, pattern, pattern)
+	if needle := strings.TrimSpace(params.Search); needle != "" {
+		cond, args := clientSearchCond(caseVariants(needle))
+		where(cond, args...)
 	}
 	if protocols := parseCSVStrings(params.Protocol); len(protocols) > 0 {
 		where("EXISTS (SELECT 1 FROM client_inbounds ci JOIN inbounds ib ON ib.id = ci.inbound_id"+
@@ -264,7 +297,8 @@ func (q clientQuery) applyParams(tx *gorm.DB, params ClientPageParams, onlines [
 		where("TRIM(COALESCE(c.comment, '')) = ''")
 	}
 	if groups := parseCSVStrings(params.Group); len(groups) > 0 {
-		where("LOWER(TRIM(COALESCE(c.group_name, ''))) IN ?", groups)
+		// The raw names cover non-ASCII capitals, which SQLite's LOWER() leaves alone.
+		where("(LOWER(TRIM(COALESCE(c.group_name, ''))) IN ? OR TRIM(COALESCE(c.group_name, '')) IN ?)", groups, groupVariants(params.Group))
 	}
 	return tx, narrowed
 }
@@ -662,6 +696,21 @@ func parseCSVStrings(raw string) []string {
 	return out
 }
 
+// groupVariants is every case spelling of each requested group name.
+func groupVariants(raw string) []string {
+	var out []string
+	for _, p := range strings.Split(raw, ",") {
+		if p = strings.TrimSpace(p); p != "" {
+			for _, v := range caseVariants(p) {
+				if !slices.Contains(out, v) {
+					out = append(out, v)
+				}
+			}
+		}
+	}
+	return out
+}
+
 // parseCSVInts is parseCSVStrings for positive integer IDs; non-numeric or
 // non-positive entries are silently dropped.
 func parseCSVInts(raw string) []int {

+ 33 - 0
internal/web/service/client_paging_test.go

@@ -635,3 +635,36 @@ func TestListPagedEmptyPanel(t *testing.T) {
 		t.Fatal("groups = nil, want an empty list so the filter drawer renders")
 	}
 }
+
+// SQLite's LOWER() folds ASCII only, so non-ASCII capitals must still match.
+func TestListPagedNonASCIICase(t *testing.T) {
+	svc, inboundSvc, settingSvc := setupPagingServices(t)
+	rec := model.ClientRecord{Email: "lima@x", Comment: "Привет", Group: "Тест", Enable: true}
+	if err := database.GetDB().Create(&rec).Error; err != nil {
+		t.Fatalf("create client: %v", err)
+	}
+	for name, params := range map[string]ClientPageParams{
+		"group":             {PageSize: 50, Group: "Тест"},
+		"group other case":  {PageSize: 50, Group: "ТЕСТ"},
+		"search":            {PageSize: 50, Search: "Привет"},
+		"search lower case": {PageSize: 50, Search: "привет"},
+		"search upper case": {PageSize: 50, Search: "ПРИВЕТ"},
+	} {
+		resp, err := svc.ListPaged(inboundSvc, settingSvc, params)
+		if err != nil {
+			t.Fatalf("%s: ListPaged: %v", name, err)
+		}
+		if got := pagedEmails(resp.Items); !slices.Equal(got, []string{"lima@x"}) {
+			t.Fatalf("%s: emails = %v, want [lima@x]", name, got)
+		}
+	}
+}
+
+func TestCaseVariants(t *testing.T) {
+	if got, want := caseVariants("пРИвет"), []string{"привет", "пРИвет", "ПРИВЕТ", "Привет"}; !slices.Equal(got, want) {
+		t.Fatalf("caseVariants = %v, want %v", got, want)
+	}
+	if got := caseVariants("abc"); !slices.Equal(got, []string{"abc", "ABC", "Abc"}) {
+		t.Fatalf("ASCII variants = %v", got)
+	}
+}