usernames: report reserved names as taken in the check paths too

account.checkUsername / channels.checkUsername / bots.checkUsername said a
reserved name was available and only updateUsername rejected it. Add the
blocklist check to peerUsernameAvailable (covers account + channel, both
backends) and to bots.Service.CheckUsername, so the client shows "username is
taken" immediately.
This commit is contained in:
Astra 2026-09-09 21:41:53 +01:00
parent 60537342d3
commit f1c24e483c
7 changed files with 79 additions and 5 deletions

View file

@ -121,6 +121,7 @@ type Service struct {
messages store.MessageStore
blocker blockChecker
channels publicChannelUsernameResolver
reserved reservedUsernameChecker
stickers stickerSetCreator
installer userStickerSetInstaller
aiChat aiChatGenerator
@ -199,6 +200,21 @@ func WithBotAvatarStore(a botAvatarStore) Option {
// WithPublicChannelUsernameResolver 注入公开频道 username 查询能力,用于 bot
// username 预检,避免 bot 与 public channel 产生同名可见入口。
// reservedUsernameChecker reports whether a name is on the operator blocklist.
type reservedUsernameChecker interface {
IsReserved(ctx context.Context, usernameLower string) (bool, error)
}
// WithReservedUsernames wires the operator username blocklist so CheckUsername
// reports a reserved bot name as taken instead of available.
func WithReservedUsernames(c reservedUsernameChecker) Option {
return func(s *Service) {
if c != nil {
s.reserved = c
}
}
}
func WithPublicChannelUsernameResolver(c publicChannelUsernameResolver) Option {
return func(s *Service) {
if c != nil {
@ -533,6 +549,13 @@ func (s *Service) CheckUsername(ctx context.Context, ownerUserID int64, username
if !domain.ValidBotUsername(username) {
return false, domain.ErrBotUsernameInvalid
}
if s.reserved != nil {
if r, err := s.reserved.IsReserved(ctx, strings.ToLower(username)); err != nil {
return false, err
} else if r {
return false, nil
}
}
if _, found, err := s.users.ByUsername(ctx, username); err != nil {
return false, err
} else if found {

View file

@ -128,6 +128,9 @@ func (s *ChannelStore) CheckUsername(_ context.Context, userID, channelID int64,
return false, err
}
usernameLower := strings.ToLower(strings.TrimSpace(strings.TrimPrefix(username, "@")))
if s.usernameRegistry != nil && s.usernameRegistry.nameReserved(usernameLower) {
return false, nil
}
for id, channel := range s.channels {
if channel.Deleted || channel.Username == "" {
continue

View file

@ -67,8 +67,10 @@ func (s *CollectibleUsernameStore) WithReservedUsernames(reserved *ReservedUsern
return s
}
func (s *CollectibleUsernameStore) nameReservedLocked(usernameLower string) bool {
if s.reserved == nil {
// nameReserved reports whether a name is on the operator blocklist. It touches
// only s.reserved (its own lock), so it is safe from any context.
func (s *CollectibleUsernameStore) nameReserved(usernameLower string) bool {
if s == nil || s.reserved == nil {
return false
}
r, _ := s.reserved.IsReserved(context.Background(), usernameLower)
@ -119,7 +121,7 @@ func (s *CollectibleUsernameStore) SetEditableUsername(_ context.Context, peer d
return false, domain.ErrUsernameInvalid
}
key := strings.ToLower(username)
if s.nameReservedLocked(key) {
if s.nameReserved(key) {
return false, domain.ErrUsernameOccupied
}
if existing, ok := s.registry[key]; ok {
@ -334,7 +336,7 @@ func (s *CollectibleUsernameStore) MintCollectibleUsername(_ context.Context, re
if _, ok := s.registry[key]; ok {
return domain.CollectibleUsername{}, false, domain.ErrUsernameOccupied
}
if s.nameReservedLocked(key) {
if s.nameReserved(key) {
return domain.CollectibleUsername{}, false, domain.ErrUsernameOccupied
}
now := time.Now().UTC()

View file

@ -0,0 +1,37 @@
package memory
import (
"context"
"testing"
"telesrv/internal/domain"
)
func TestCheckUsernameReportsReservedAsTaken(t *testing.T) {
ctx := context.Background()
reserved := NewReservedUsernameStore()
if _, err := reserved.ReserveUsername(ctx, "support", "official", "ops"); err != nil {
t.Fatalf("seed reserve: %v", err)
}
registry := NewCollectibleUsernameStore().WithReservedUsernames(reserved)
users := NewUserStore()
users.AttachUsernameRegistry(registry)
u, _ := users.Create(ctx, domain.User{AccessHash: 1, Phone: "15550001000", FirstName: "A"})
if ok, err := users.CheckUsername(ctx, u.ID, "support"); err != nil || ok {
t.Fatalf("CheckUsername(reserved) = %v, %v; want false, nil", ok, err)
}
if ok, err := users.CheckUsername(ctx, u.ID, "freename"); err != nil || !ok {
t.Fatalf("CheckUsername(free) = %v, %v; want true, nil", ok, err)
}
channels := NewChannelStore()
channels.AttachUsernameRegistry(registry)
created, err := channels.CreateChannel(ctx, domain.CreateChannelRequest{CreatorUserID: u.ID, Title: "C", Megagroup: true, Date: 1})
if err != nil {
t.Fatalf("create channel: %v", err)
}
if ok, err := channels.CheckUsername(ctx, u.ID, created.Channel.ID, "support"); err != nil || ok {
t.Fatalf("channel CheckUsername(reserved) = %v, %v; want false, nil", ok, err)
}
}

View file

@ -160,6 +160,9 @@ func (s *UserStore) CheckUsername(_ context.Context, userID int64, username stri
if username == "" {
return true, nil
}
if s.usernameRegistry != nil && s.usernameRegistry.nameReserved(username) {
return false, nil
}
s.mu.RLock()
defer s.mu.RUnlock()
for id, u := range s.byID {

View file

@ -81,6 +81,11 @@ func usernameReservedTx(ctx context.Context, db sqlcgen.DBTX, usernameLower stri
}
func peerUsernameAvailable(ctx context.Context, db sqlcgen.DBTX, usernameLower, peerType string, peerID int64) (bool, error) {
if reserved, err := usernameReservedTx(ctx, db, usernameLower); err != nil {
return false, err
} else if reserved {
return false, nil
}
owner, found, err := getPeerUsernameOwner(ctx, db, usernameLower, false)
if err != nil || !found {
return !found, err