diff --git a/docs/v3-api.yaml b/docs/v3-api.yaml index bfb0f50f6..6ed7eb872 100644 --- a/docs/v3-api.yaml +++ b/docs/v3-api.yaml @@ -6068,8 +6068,10 @@ components: type: string description: FCMのデバイストークン example: "bk3RNwTe3H0:CI2k_HHwgIpoDKCIZvvDMExUdFQ3P1" - required: - - token + fid: + type: string + description: FCMのInstallation ID + example: "cA7kP2mV9xQn4Lw8RjT0sB" PostUserRequest: title: PostUserRequest type: object diff --git a/go.mod b/go.mod index 9231bdf91..c18bc005d 100644 --- a/go.mod +++ b/go.mod @@ -6,7 +6,7 @@ toolchain go1.26.5 require ( cloud.google.com/go/profiler v0.6.0 - firebase.google.com/go/v4 v4.20.0 + firebase.google.com/go/v4 v4.21.0 github.com/MicahParks/jwkset v0.11.3 github.com/NYTimes/gziphandler v1.1.1 github.com/aws/aws-sdk-go-v2 v1.43.3 diff --git a/go.sum b/go.sum index b0660be02..c42904bf1 100644 --- a/go.sum +++ b/go.sum @@ -68,8 +68,8 @@ dario.cat/mergo v1.0.2/go.mod h1:E/hbnu0NxMFBjpMIE34DRGLWqDy0g5FuKDhCb31ngxA= dmitri.shuralyov.com/gpu/mtl v0.0.0-20190408044501-666a987793e9/go.mod h1:H6x//7gZCb22OMCxBHrMx7a5I7Hp++hsVxbQ4BYO7hU= filippo.io/edwards25519 v1.2.0 h1:crnVqOiS4jqYleHd9vaKZ+HKtHfllngJIiOpNpoJsjo= filippo.io/edwards25519 v1.2.0/go.mod h1:xzAOLCNug/yB62zG1bQ8uziwrIqIuxhctzJT18Q77mc= -firebase.google.com/go/v4 v4.20.0 h1:ighpjeAC45rY/95cUQ+ojIKlKcTnz2YC0ldam56z2YU= -firebase.google.com/go/v4 v4.20.0/go.mod h1:hqhkQtZkThGH42TnaYi7A8EFR1E0FEuB5oHvJ1Q57t8= +firebase.google.com/go/v4 v4.21.0 h1:HBZV4jrLtFYj8EwWyqEZOuRLfkfkV2bpnfyyXHOhPxY= +firebase.google.com/go/v4 v4.21.0/go.mod h1:CDumIdA5oTiyDpLNVcQoW8ZrB5CTgyE2D45DuENIABg= github.com/Azure/go-ansiterm v0.0.0-20250102033503-faa5f7b0171c h1:udKWzYgxTojEKWjV8V+WSxDXJ4NFATAsZjh8iIbsQIg= github.com/Azure/go-ansiterm v0.0.0-20250102033503-faa5f7b0171c/go.mod h1:xomTg63KZ2rFqZQzSB4Vz2SUXa1BpHTVz9L5PTmPC4E= github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= diff --git a/migration/current.go b/migration/current.go index 3d40bb137..d570e3ce6 100644 --- a/migration/current.go +++ b/migration/current.go @@ -11,49 +11,50 @@ import ( // 新たなマイグレーションを行う場合は、この配列の末尾に必ず追加すること func Migrations() []*gormigrate.Migration { return []*gormigrate.Migration{ - v1(), // インデックスidx_messages_deleted_atの削除とidx_messages_channel_id_deleted_at_created_atの追加 - v2(), // RBAC周りのリフォーム - v3(), // チャンネルイベント履歴 - v4(), // Webhook, Bot外部キー - v5(), // Mute, 旧Clip削除 - v6(), // ユーザーグループ拡張 - v7(), // ファイルメタ拡張 - v8(), // チャンネル購読拡張 - v9(), // ユーザーテーブル拡張 - v10(), // パーミッション周りの調整 - v11(), // クリップ機能の追加 - v12(), // カスタムスタンプパレットの追加 - v13(), // パーミッション調整・インデックス付与 - v14(), // パーミッション不足修正 - v15(), // 外部ログイン機能追加 - v16(), // パーミッション修正 - v17(), // ユーザーホームチャンネル - v18(), // インデックス追加 - v19(), // httpセッション管理テーブル変更 - v20(), // パーミッション周りの調整 - v21(), // OGPキャッシュ追加 - v22(), // BOTへのWebRTCパーミッションの付与 - v23(), // 複合インデックス追加 - v24(), // ユーザー設定追加 - v25(), // FileMetaにIsAnimatedImageを追加 - v26(), // FileMetaからThumbnail情報を分離 - v27(), // Gorm v2移行: FKの追加、FKのリネーム、一部フィールドのデータ型変更、idx_messages_channel_idの削除 - v28(), // ユーザーグループにアイコンを追加 - v29(), // BotにModeを追加、WebSocket Modeを追加 - v30(), // bot_event_logsにresultを追加 - v31(), // お気に入りスタンプパーミッション削除(削除忘れ) - v32(), // ユーザーの表示名上限を32文字に - v33(), // 未読テーブルにチャンネルIDカラムを追加 / インデックス類の更新 / 不要なレコードの削除 - v34(), // 未読テーブルのcreated_atカラムをメッセージテーブルを元に更新 / カラム名を変更 - v35(), // OIDC実装のため、openid, profileロール、get_oidc_userinfo権限を追加 - v36(), // delete_my_stampパーミッションを追加 - v37(), // サウンドボードアイテム追加 - v38(), // v37で作ったサウンドボードアイテムのテーブル名変更 - v39(), // OAuth Client Credentials Grantの対応のため、clientロールを追加 - v40(), // delete_my_stampパーミッションを削除 - v41(), // ユーザーグループ名受付規則変更に伴う既存ユーザーグループ名の更新 - v42(), // get_my_stamp_recommendationsパーミッションの追加とmessages_stampsテーブルへの (user_id, updated_at) の複合インデックスの追加 - v43(), // messages_stampsテーブルのインデックス (user_id, updated_at) を (user_id, updated_at, stamp_id) に変更 + v1(), // インデックスidx_messages_deleted_atの削除とidx_messages_channel_id_deleted_at_created_atの追加 + v2(), // RBAC周りのリフォーム + v3(), // チャンネルイベント履歴 + v4(), // Webhook, Bot外部キー + v5(), // Mute, 旧Clip削除 + v6(), // ユーザーグループ拡張 + v7(), // ファイルメタ拡張 + v8(), // チャンネル購読拡張 + v9(), // ユーザーテーブル拡張 + v10(), // パーミッション周りの調整 + v11(), // クリップ機能の追加 + v12(), // カスタムスタンプパレットの追加 + v13(), // パーミッション調整・インデックス付与 + v14(), // パーミッション不足修正 + v15(), // 外部ログイン機能追加 + v16(), // パーミッション修正 + v17(), // ユーザーホームチャンネル + v18(), // インデックス追加 + v19(), // httpセッション管理テーブル変更 + v20(), // パーミッション周りの調整 + v21(), // OGPキャッシュ追加 + v22(), // BOTへのWebRTCパーミッションの付与 + v23(), // 複合インデックス追加 + v24(), // ユーザー設定追加 + v25(), // FileMetaにIsAnimatedImageを追加 + v26(), // FileMetaからThumbnail情報を分離 + v27(), // Gorm v2移行: FKの追加、FKのリネーム、一部フィールドのデータ型変更、idx_messages_channel_idの削除 + v28(), // ユーザーグループにアイコンを追加 + v29(), // BotにModeを追加、WebSocket Modeを追加 + v30(), // bot_event_logsにresultを追加 + v31(), // お気に入りスタンプパーミッション削除(削除忘れ) + v32(), // ユーザーの表示名上限を32文字に + v33(), // 未読テーブルにチャンネルIDカラムを追加 / インデックス類の更新 / 不要なレコードの削除 + v34(), // 未読テーブルのcreated_atカラムをメッセージテーブルを元に更新 / カラム名を変更 + v35(), // OIDC実装のため、openid, profileロール、get_oidc_userinfo権限を追加 + v36(), // delete_my_stampパーミッションを追加 + v37(), // サウンドボードアイテム追加 + v38(), // v37で作ったサウンドボードアイテムのテーブル名変更 + v39(), // OAuth Client Credentials Grantの対応のため、clientロールを追加 + v40(), // delete_my_stampパーミッションを削除 + v41(), // ユーザーグループ名受付規則変更に伴う既存ユーザーグループ名の更新 + v42(), // get_my_stamp_recommendationsパーミッションの追加とmessages_stampsテーブルへの (user_id, updated_at) の複合インデックスの追加 + v43(), // messages_stampsテーブルのインデックス (user_id, updated_at) を (user_id, updated_at, stamp_id) に変更 + v43_1(), // Firebase Installation ID に対応するため、TokenType カラムを追加 } } diff --git a/migration/v43-1.go b/migration/v43-1.go new file mode 100644 index 000000000..0c51336d5 --- /dev/null +++ b/migration/v43-1.go @@ -0,0 +1,36 @@ +package migration + +import ( + "time" + + "github.com/go-gormigrate/gormigrate/v2" + "github.com/gofrs/uuid" + "github.com/traPtitech/traQ/model" + "gorm.io/gorm" +) + +// v43-1 Added TokenType column in order to handle the Firebase Installation ID +func v43_1() *gormigrate.Migration { + return &gormigrate.Migration{ + ID: "43-1", + Migrate: func(db *gorm.DB) error { + if err := db.Exec("ALTER TABLE devices DROP PRIMARY KEY, ADD COLUMN token_type VARCHAR(16) DEFAULT 'token' NOT NULL AFTER token, ADD PRIMARY KEY (token, token_type)").Error; err != nil { + return err + } + return db.AutoMigrate(&v43_1Device{}) + }, + } +} + +type v43_1Device struct { + Token string `gorm:"type:varchar(190);not null;primaryKey"` + TokenType model.DeviceTokenType `gorm:"type:varchar(16);not null;primaryKey"` + UserID uuid.UUID `gorm:"type:char(36);not null;index"` + CreatedAt time.Time `gorm:"precision:6"` + + User *model.User `gorm:"constraint:devices_user_id_users_id_foreign,OnUpdate:CASCADE,OnDelete:CASCADE"` +} + +func (*v43_1Device) TableName() string { + return "devices" +} diff --git a/model/devices.go b/model/devices.go index 5e6f8a4f6..cd81f9e17 100644 --- a/model/devices.go +++ b/model/devices.go @@ -6,15 +6,23 @@ import ( "github.com/gofrs/uuid" ) +type DeviceTokenType string + // Device 通知デバイスの構造体 type Device struct { - Token string `gorm:"type:varchar(190);not null;primaryKey"` - UserID uuid.UUID `gorm:"type:char(36);not null;index"` - CreatedAt time.Time `gorm:"precision:6"` + Token string `gorm:"type:varchar(190);not null;primaryKey"` + TokenType DeviceTokenType `gorm:"type:varchar(16);not null;primaryKey"` + UserID uuid.UUID `gorm:"type:char(36);not null;index"` + CreatedAt time.Time `gorm:"precision:6"` User *User `gorm:"constraint:devices_user_id_users_id_foreign,OnUpdate:CASCADE,OnDelete:CASCADE"` } +const ( + DeviceTokenTypeToken DeviceTokenType = "token" // レガシーな Token + DeviceTokenTypeFID DeviceTokenType = "fid" // Firebase Installation ID +) + // TableName Device構造体のテーブル名 func (*Device) TableName() string { return "devices" diff --git a/repository/device.go b/repository/device.go index b86a7d6d0..2d0c61608 100644 --- a/repository/device.go +++ b/repository/device.go @@ -5,9 +5,20 @@ import ( "github.com/gofrs/uuid" + "github.com/traPtitech/traQ/model" + "github.com/traPtitech/traQ/utils/optional" "github.com/traPtitech/traQ/utils/set" ) +type RegisterDeviceArgs struct { + Token optional.Of[string] + FID optional.Of[string] +} +type TokenEntry struct { + Token string + TokenType model.DeviceTokenType +} + // DeviceRepository FCMデバイスリポジトリ type DeviceRepository interface { // RegisterDevice FCMデバイスを登録します @@ -17,12 +28,12 @@ type DeviceRepository interface { // tokenが空文字列の場合、ArgumentErrorを返します。 // 登録しようとしたトークンが既に他のユーザーと関連づけられていた場合はArgumentErrorを返します。 // DBによるエラーを返すことがあります。 - RegisterDevice(ctx context.Context, userID uuid.UUID, token string) error + RegisterDevice(ctx context.Context, userID uuid.UUID, args RegisterDeviceArgs) error // GetDeviceTokens 指定したユーザーの全デバイストークンを取得します // // 成功した場合、デバイストークンの配列とnilを返します。 // DBによるエラーを返すことがあります。 - GetDeviceTokens(ctx context.Context, userIDs set.UUID) (map[uuid.UUID][]string, error) + GetDeviceTokens(ctx context.Context, userIDs set.UUID) (map[uuid.UUID][]TokenEntry, error) // DeleteDeviceTokens FCMデバイスの登録を解除します // // 成功した、或いは既に登録解除されていた場合にnilを返します。 diff --git a/repository/gorm/device.go b/repository/gorm/device.go index 93a008629..9f44f57b5 100644 --- a/repository/gorm/device.go +++ b/repository/gorm/device.go @@ -12,17 +12,32 @@ import ( ) // RegisterDevice implements DeviceRepository interface. -func (repo *Repository) RegisterDevice(ctx context.Context, userID uuid.UUID, token string) error { +func (repo *Repository) RegisterDevice(ctx context.Context, userID uuid.UUID, args repository.RegisterDeviceArgs) error { if userID == uuid.Nil { return repository.ErrNilID } - if len(token) == 0 { - return repository.ArgError("Token", "token is empty") + + token := args.FID.V + tokenType := model.DeviceTokenTypeFID + if !args.FID.Valid { + token = args.Token.V + tokenType = model.DeviceTokenTypeToken + if !args.Token.Valid { + return repository.ArgError("Token, FID", "token and fid are empty") + } } err := repo.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { var d model.Device - if err := tx.First(&d, &model.Device{Token: token}).Error; err == nil { + if len(args.Token.ValueOrZero()) != 0 && len(args.FID.ValueOrZero()) != 0 { + if err := tx.Delete(&model.Device{Token: args.Token.ValueOrZero(), TokenType: model.DeviceTokenTypeToken}).Error; err != nil { + if err != gorm.ErrRecordNotFound { + return err + } + } + } + + if err := tx.First(&d, &model.Device{Token: token, TokenType: tokenType}).Error; err == nil { if d.UserID != userID { return repository.ArgError("Token", "the Token has already been associated with other user") } @@ -32,23 +47,24 @@ func (repo *Repository) RegisterDevice(ctx context.Context, userID uuid.UUID, to } return tx.Create(&model.Device{ - Token: token, - UserID: userID, + Token: token, + TokenType: tokenType, + UserID: userID, }).Error }) return err } // GetDeviceTokens implements DeviceRepository interface. -func (repo *Repository) GetDeviceTokens(ctx context.Context, userIDs set.UUID) (tokens map[uuid.UUID][]string, err error) { +func (repo *Repository) GetDeviceTokens(ctx context.Context, userIDs set.UUID) (tokens map[uuid.UUID][]repository.TokenEntry, err error) { var tmp []*model.Device if err := repo.db.WithContext(ctx).Where("user_id IN (?)", userIDs.StringArray()).Find(&tmp).Error; err != nil { return nil, err } - tokens = make(map[uuid.UUID][]string, len(userIDs)) + tokens = make(map[uuid.UUID][]repository.TokenEntry, len(userIDs)) for _, device := range tmp { - tokens[device.UserID] = append(tokens[device.UserID], device.Token) + tokens[device.UserID] = append(tokens[device.UserID], repository.TokenEntry{Token: device.Token, TokenType: device.TokenType}) } return tokens, nil } diff --git a/repository/gorm/device_test.go b/repository/gorm/device_test.go index f5e666088..2a0a90ffa 100644 --- a/repository/gorm/device_test.go +++ b/repository/gorm/device_test.go @@ -8,6 +8,8 @@ import ( "github.com/stretchr/testify/assert" "github.com/traPtitech/traQ/model" + "github.com/traPtitech/traQ/repository" + "github.com/traPtitech/traQ/utils/optional" random2 "github.com/traPtitech/traQ/utils/random" "github.com/traPtitech/traQ/utils/set" ) @@ -20,25 +22,67 @@ func TestRepositoryImpl_RegisterDevice(t *testing.T) { id2 := mustMakeUser(t, repo, rand, false).GetID() id3 := mustMakeUser(t, repo, rand, true).GetID() id4 := mustMakeUser(t, repo, rand, true).GetID() + id5 := mustMakeUser(t, repo, rand, false).GetID() + id6 := mustMakeUser(t, repo, rand, false).GetID() + id7 := mustMakeUser(t, repo, rand, true).GetID() + id8 := mustMakeUser(t, repo, rand, true).GetID() + id9 := mustMakeUser(t, repo, rand, false).GetID() + id10 := mustMakeUser(t, repo, rand, false).GetID() + id11 := mustMakeUser(t, repo, rand, true).GetID() + id12 := mustMakeUser(t, repo, rand, true).GetID() token1 := random2.AlphaNumeric(20) token2 := random2.AlphaNumeric(20) token3 := random2.AlphaNumeric(20) token4 := random2.AlphaNumeric(20) + token5 := random2.AlphaNumeric(20) + token6 := random2.AlphaNumeric(20) + token7 := random2.AlphaNumeric(20) + token8 := random2.AlphaNumeric(20) + + fid1 := random2.AlphaNumeric(22) + fid2 := random2.AlphaNumeric(22) + fid3 := random2.AlphaNumeric(22) + fid4 := random2.AlphaNumeric(22) + fid5 := random2.AlphaNumeric(22) + fid6 := random2.AlphaNumeric(22) + fid7 := random2.AlphaNumeric(22) + fid8 := random2.AlphaNumeric(22) cases := []struct { user uuid.UUID - token string + token repository.RegisterDeviceArgs error bool }{ - {id1, token1, false}, - {id2, token2, false}, - {id2, token2, false}, - {id3, token3, false}, - {id4, token4, false}, - {id1, token2, true}, - {uuid.Nil, token2, true}, - {id1, "", true}, + {id1, repository.RegisterDeviceArgs{Token: optional.New(token1, true)}, false}, + {id2, repository.RegisterDeviceArgs{Token: optional.New(token2, true)}, false}, + {id2, repository.RegisterDeviceArgs{Token: optional.New(token2, true)}, false}, + {id3, repository.RegisterDeviceArgs{Token: optional.New(token3, true)}, false}, + {id4, repository.RegisterDeviceArgs{Token: optional.New(token4, true)}, false}, + {id1, repository.RegisterDeviceArgs{Token: optional.New(token2, true)}, true}, + {uuid.Nil, repository.RegisterDeviceArgs{Token: optional.New(token2, true)}, true}, + {id1, repository.RegisterDeviceArgs{Token: optional.New("", true)}, true}, + {id1, repository.RegisterDeviceArgs{Token: optional.New("", false)}, true}, + + {id5, repository.RegisterDeviceArgs{FID: optional.New(fid1, true)}, false}, + {id6, repository.RegisterDeviceArgs{FID: optional.New(fid2, true)}, false}, + {id6, repository.RegisterDeviceArgs{FID: optional.New(fid2, true)}, false}, + {id7, repository.RegisterDeviceArgs{FID: optional.New(fid3, true)}, false}, + {id8, repository.RegisterDeviceArgs{FID: optional.New(fid4, true)}, false}, + {id5, repository.RegisterDeviceArgs{FID: optional.New(fid2, true)}, true}, + {uuid.Nil, repository.RegisterDeviceArgs{FID: optional.New(fid2, true)}, true}, + {id5, repository.RegisterDeviceArgs{FID: optional.New("", true)}, true}, + {id5, repository.RegisterDeviceArgs{FID: optional.New("", false)}, true}, + + {id9, repository.RegisterDeviceArgs{Token: optional.New(token5, true), FID: optional.New(fid5, true)}, false}, + {id10, repository.RegisterDeviceArgs{Token: optional.New(token6, true), FID: optional.New(fid6, true)}, false}, + {id10, repository.RegisterDeviceArgs{Token: optional.New(token6, true), FID: optional.New(fid6, true)}, false}, + {id11, repository.RegisterDeviceArgs{Token: optional.New(token7, true), FID: optional.New(fid7, true)}, false}, + {id12, repository.RegisterDeviceArgs{Token: optional.New(token8, true), FID: optional.New(fid8, true)}, false}, + {id9, repository.RegisterDeviceArgs{Token: optional.New(token6, true), FID: optional.New(fid6, true)}, true}, + {uuid.Nil, repository.RegisterDeviceArgs{Token: optional.New(token6, true), FID: optional.New(fid6, true)}, true}, + {id9, repository.RegisterDeviceArgs{Token: optional.New("", true), FID: optional.New("", true)}, true}, + {id9, repository.RegisterDeviceArgs{Token: optional.New("", false), FID: optional.New("", false)}, true}, } for _, v := range cases { @@ -50,7 +94,7 @@ func TestRepositoryImpl_RegisterDevice(t *testing.T) { } } - assert.EqualValues(4, count(t, getDB(repo).Model(model.Device{}).Where("user_id IN (?, ? ,? , ?)", id1, id2, id3, id4))) + assert.EqualValues(12, count(t, getDB(repo).Model(model.Device{}).Where("user_id IN (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", id1, id2, id3, id4, id5, id6, id7, id8, id9, id10, id11, id12))) } func TestRepositoryImpl_DeleteDeviceTokens(t *testing.T) { @@ -70,19 +114,19 @@ func TestRepositoryImpl_DeleteDeviceTokens(t *testing.T) { token6 := random2.AlphaNumeric(20) token7 := random2.AlphaNumeric(20) - err := repo.RegisterDevice(context.TODO(), id1, token1) + err := repo.RegisterDevice(context.TODO(), id1, repository.RegisterDeviceArgs{Token: optional.New(token1, true)}) require.NoError(err) - err = repo.RegisterDevice(context.TODO(), id2, token2) + err = repo.RegisterDevice(context.TODO(), id2, repository.RegisterDeviceArgs{Token: optional.New(token2, true)}) require.NoError(err) - err = repo.RegisterDevice(context.TODO(), id1, token3) + err = repo.RegisterDevice(context.TODO(), id1, repository.RegisterDeviceArgs{Token: optional.New(token3, true)}) require.NoError(err) - err = repo.RegisterDevice(context.TODO(), id1, token4) + err = repo.RegisterDevice(context.TODO(), id1, repository.RegisterDeviceArgs{Token: optional.New(token4, true)}) require.NoError(err) - err = repo.RegisterDevice(context.TODO(), id3, token5) + err = repo.RegisterDevice(context.TODO(), id3, repository.RegisterDeviceArgs{Token: optional.New(token5, true)}) require.NoError(err) - err = repo.RegisterDevice(context.TODO(), id4, token6) + err = repo.RegisterDevice(context.TODO(), id4, repository.RegisterDeviceArgs{Token: optional.New(token6, true)}) require.NoError(err) - err = repo.RegisterDevice(context.TODO(), id4, token7) + err = repo.RegisterDevice(context.TODO(), id4, repository.RegisterDeviceArgs{Token: optional.New(token7, true)}) require.NoError(err) cases := []struct { @@ -115,13 +159,13 @@ func TestRepositoryImpl_GetDeviceTokens(t *testing.T) { token3 := random2.AlphaNumeric(20) token4 := random2.AlphaNumeric(20) - err := repo.RegisterDevice(context.TODO(), id1, token1) + err := repo.RegisterDevice(context.TODO(), id1, repository.RegisterDeviceArgs{Token: optional.New(token1, true)}) require.NoError(err) - err = repo.RegisterDevice(context.TODO(), id2, token2) + err = repo.RegisterDevice(context.TODO(), id2, repository.RegisterDeviceArgs{Token: optional.New(token2, true)}) require.NoError(err) - err = repo.RegisterDevice(context.TODO(), id1, token3) + err = repo.RegisterDevice(context.TODO(), id1, repository.RegisterDeviceArgs{Token: optional.New(token3, true)}) require.NoError(err) - err = repo.RegisterDevice(context.TODO(), id3, token4) + err = repo.RegisterDevice(context.TODO(), id3, repository.RegisterDeviceArgs{Token: optional.New(token4, true)}) require.NoError(err) cases := []struct { diff --git a/router/v3/users.go b/router/v3/users.go index 251872da9..a516d7cf7 100644 --- a/router/v3/users.go +++ b/router/v3/users.go @@ -327,12 +327,14 @@ func (h *Handlers) GetMyStampRecommendations(c *echo.Context) error { // PostMyFCMDeviceRequest POST /users/me/fcm-device リクエストボディ type PostMyFCMDeviceRequest struct { - Token string `json:"token"` + Token optional.Of[string] `json:"token"` + FID optional.Of[string] `json:"fid"` } func (r PostMyFCMDeviceRequest) Validate() error { return vd.ValidateStruct(&r, - vd.Field(&r.Token, vd.Required, vd.RuneLength(1, 190)), + vd.Field(&r.Token, vd.RuneLength(1, 190), vd.When(!r.FID.Valid, vd.Required)), + vd.Field(&r.FID, vd.RuneLength(22, 22), vd.When(!r.Token.Valid, vd.Required)), ) } @@ -343,8 +345,13 @@ func (h *Handlers) PostMyFCMDevice(c *echo.Context) error { return err } + args := repository.RegisterDeviceArgs{ + Token: req.Token, + FID: req.FID, + } + userID := getRequestUserID(c) - if err := h.Repo.RegisterDevice(c.Request().Context(), userID, req.Token); err != nil { + if err := h.Repo.RegisterDevice(c.Request().Context(), userID, args); err != nil { switch { case repository.IsArgError(err): return herror.BadRequest(err) diff --git a/router/v3/users_test.go b/router/v3/users_test.go index 4c958f6cb..1790be9e5 100644 --- a/router/v3/users_test.go +++ b/router/v3/users_test.go @@ -662,7 +662,8 @@ func TestPostMyFCMDeviceRequest_Validate(t *testing.T) { t.Parallel() type fields struct { - Token string + Token optional.Of[string] + FID optional.Of[string] } tests := []struct { name string @@ -676,7 +677,7 @@ func TestPostMyFCMDeviceRequest_Validate(t *testing.T) { }, { "success", - fields{Token: "dummy:token"}, + fields{Token: optional.From("dummy:token")}, false, }, } @@ -701,37 +702,67 @@ func TestHandlers_PostMyFCMDevice(t *testing.T) { s := env.S(t, user.GetID()) t.Run("not logged in", func(t *testing.T) { - t.Parallel() e := env.R(t) e.POST(path). - WithJSON(&PostMyFCMDeviceRequest{Token: "dummy:token"}). + WithJSON(&PostMyFCMDeviceRequest{Token: optional.From("dummy:token")}). Expect(). Status(http.StatusUnauthorized) }) t.Run("bad request (empty token)", func(t *testing.T) { - t.Parallel() e := env.R(t) e.POST(path). WithCookie(session.CookieName, s). - WithJSON(&PostMyFCMDeviceRequest{Token: ""}). + WithJSON(&PostMyFCMDeviceRequest{Token: optional.From("")}). Expect(). Status(http.StatusBadRequest) }) t.Run("success", func(t *testing.T) { - t.Parallel() e := env.R(t) e.POST(path). WithCookie(session.CookieName, s). - WithJSON(&PostMyFCMDeviceRequest{Token: "dummy:token"}). + WithJSON(&PostMyFCMDeviceRequest{Token: optional.From("dummy:token")}). + Expect(). + Status(http.StatusNoContent) + + tokens, err := env.Repository.GetDeviceTokens(context.TODO(), set.UUID{user.GetID(): {}}) + require.NoError(t, err) + if assert.Len(t, tokens, 1) { + assert.ElementsMatch(t, tokens[user.GetID()], []repository.TokenEntry{{Token: "dummy:token", TokenType: model.DeviceTokenTypeToken}}) + } + }) + + t.Run("success", func(t *testing.T) { + e := env.R(t) + e.POST(path). + WithCookie(session.CookieName, s). + WithJSON(&PostMyFCMDeviceRequest{FID: optional.From("dummy:token23456789012")}). + Expect(). + Status(http.StatusNoContent) + + tokens, err := env.Repository.GetDeviceTokens(context.TODO(), set.UUID{user.GetID(): {}}) + require.NoError(t, err) + if assert.Len(t, tokens, 1) { + assert.ElementsMatch(t, tokens[user.GetID()], []repository.TokenEntry{ + {Token: "dummy:token", TokenType: model.DeviceTokenTypeToken}, + {Token: "dummy:token23456789012", TokenType: model.DeviceTokenTypeFID}, + }) + } + }) + + t.Run("success", func(t *testing.T) { + e := env.R(t) + e.POST(path). + WithCookie(session.CookieName, s). + WithJSON(&PostMyFCMDeviceRequest{Token: optional.From("dummy:token"), FID: optional.From("dummy:token23456789012")}). Expect(). Status(http.StatusNoContent) tokens, err := env.Repository.GetDeviceTokens(context.TODO(), set.UUID{user.GetID(): {}}) require.NoError(t, err) if assert.Len(t, tokens, 1) { - assert.ElementsMatch(t, tokens[user.GetID()], []string{"dummy:token"}) + assert.ElementsMatch(t, tokens[user.GetID()], []repository.TokenEntry{{Token: "dummy:token23456789012", TokenType: model.DeviceTokenTypeFID}}) } }) } diff --git a/service/fcm/impl.go b/service/fcm/impl.go index b1b11c478..b8745b4ad 100644 --- a/service/fcm/impl.go +++ b/service/fcm/impl.go @@ -11,6 +11,7 @@ import ( "go.uber.org/zap" "google.golang.org/api/option" + "github.com/traPtitech/traQ/model" "github.com/traPtitech/traQ/repository" "github.com/traPtitech/traQ/service/counter" "github.com/traPtitech/traQ/service/variable" @@ -100,6 +101,7 @@ func (c *clientImpl) send(targetUserIDs set.UUID, p *Payload, withUnreadCount bo Body: p.Body, } ) + // TODO: refactor if withUnreadCount { for uid, tokens := range tokensMap { unread := c.unreadCounter.Get(uid) @@ -128,13 +130,23 @@ func (c *clientImpl) send(targetUserIDs set.UUID, p *Payload, withUnreadCount bo } for _, token := range tokens { - messages = append(messages, &messaging.Message{ - Data: data, - Android: defaultAndroidConfig, - Webpush: defaultWebpushConfig, - APNS: apns, - Token: token, - }) + if token.TokenType == model.DeviceTokenTypeToken { + messages = append(messages, &messaging.Message{ + Data: data, + Android: defaultAndroidConfig, + Webpush: defaultWebpushConfig, + APNS: apns, + Token: token.Token, + }) + } else if token.TokenType == model.DeviceTokenTypeFID { + messages = append(messages, &messaging.Message{ + Data: data, + Android: defaultAndroidConfig, + Webpush: defaultWebpushConfig, + APNS: apns, + Fid: token.Token, + }) + } } } } else { @@ -162,13 +174,23 @@ func (c *clientImpl) send(targetUserIDs set.UUID, p *Payload, withUnreadCount bo for _, tokens := range tokensMap { for _, token := range tokens { - messages = append(messages, &messaging.Message{ - Data: data, - Android: defaultAndroidConfig, - Webpush: defaultWebpushConfig, - APNS: apns, - Token: token, - }) + if token.TokenType == model.DeviceTokenTypeToken { + messages = append(messages, &messaging.Message{ + Data: data, + Android: defaultAndroidConfig, + Webpush: defaultWebpushConfig, + APNS: apns, + Token: token.Token, + }) + } else if token.TokenType == model.DeviceTokenTypeFID { + messages = append(messages, &messaging.Message{ + Data: data, + Android: defaultAndroidConfig, + Webpush: defaultWebpushConfig, + APNS: apns, + Fid: token.Token, + }) + } } } }