feat(message-gateway): add user bind and unbind APIs

This commit is contained in:
ryan
2026-08-16 12:19:32 +08:00
parent 69d39d906f
commit 635c1760ad
10 changed files with 1118 additions and 0 deletions
+16
View File
@@ -0,0 +1,16 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import "errors"
var (
errCodeInvalid = errors.New("invalid or expired pairing code")
errChannelMismatch = errors.New("pairing code does not match channel")
errPlatformAlreadyBound = errors.New("this platform account is already bound")
errBindingNotFound = errors.New("binding not found")
errBindingForbidden = errors.New("cannot unbind another user's binding")
errChannelIDRequired = errors.New("channel_id is required")
errChannelDisabled = errors.New("channel is not enabled")
)
+136
View File
@@ -0,0 +1,136 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"errors"
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/gin-gonic/gin"
)
func currentUser(c *gin.Context) (*model.User, bool) {
return oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
}
// ListChannels lists enabled channels a user can bind.
// @Summary List enabled messaging channels
// @Description Returns enabled system bots the current user can pair with
// @Tags message-gateway
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]PublicChannelDTO}
// @Failure 401 {object} response.Any
// @Router /api/v1/message-gateway/channels [get]
func ListChannels(c *gin.Context) {
if user, ok := currentUser(c); !ok || user == nil {
response.AbortUnauthorized(c, "login required")
return
}
rows, err := listEnabledPublicChannels(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(rows))
}
// ListBindings lists the current user's bot bindings.
// @Summary List message gateway bindings
// @Description Returns the current user's bound messaging channels
// @Tags message-gateway
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]BindingDTO}
// @Failure 401 {object} response.Any
// @Router /api/v1/message-gateway/bindings [get]
func ListBindings(c *gin.Context) {
user, ok := currentUser(c)
if !ok || user == nil {
response.AbortUnauthorized(c, "login required")
return
}
rows, err := listUserBindings(c.Request.Context(), user.ID)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(rows))
}
// BindBinding consumes a pairing code and binds the platform identity.
// @Summary Bind a messaging channel
// @Description Binds the current user to a platform identity using a one-time pairing code
// @Tags message-gateway
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body BindRequest true "bind body"
// @Success 200 {object} response.Any{data=BindingDTO}
// @Failure 400 {object} response.Any
// @Failure 409 {object} response.Any
// @Router /api/v1/message-gateway/bindings [post]
func BindBinding(c *gin.Context) {
user, ok := currentUser(c)
if !ok || user == nil {
response.AbortUnauthorized(c, "login required")
return
}
var req BindRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
dto, err := bindChannel(c.Request.Context(), user.ID, req)
if err != nil {
if errors.Is(err, errPlatformAlreadyBound) {
response.AbortConflict(c, err.Error())
return
}
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(dto))
}
// UnbindBinding removes the current user's binding.
// @Summary Unbind a messaging channel
// @Description Removes a binding owned by the current user
// @Tags message-gateway
// @Produce json
// @Security SessionCookie
// @Param id path int true "binding id"
// @Success 200 {object} response.Any
// @Failure 403 {object} response.Any
// @Failure 404 {object} response.Any
// @Router /api/v1/message-gateway/bindings/{id} [delete]
func UnbindBinding(c *gin.Context) {
user, ok := currentUser(c)
if !ok || user == nil {
response.AbortUnauthorized(c, "login required")
return
}
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid binding id")
return
}
if err := unbindChannel(c.Request.Context(), user.ID, id); err != nil {
if errors.Is(err, errBindingNotFound) {
response.AbortNotFound(c, err.Error())
return
}
if errors.Is(err, errBindingForbidden) {
response.AbortForbidden(c, err.Error())
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
+158
View File
@@ -0,0 +1,158 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"errors"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
pkgmg "github.com/Rain-kl/Wavelet/pkg/message_gateway"
"gorm.io/gorm"
)
// BindRequest is the user bind body.
type BindRequest struct {
ChannelID string `json:"channel_id"`
Code string `json:"code"`
}
// BindingDTO is a user-facing binding row.
type BindingDTO struct {
ID uint64 `json:"id,string"`
UserID uint64 `json:"user_id,string"`
ChannelID uint64 `json:"channel_id,string"`
ChannelName string `json:"channel_name"`
ChannelType string `json:"channel_type"`
PlatformUserID string `json:"platform_user_id"`
CreatedAt time.Time `json:"created_at"`
}
func bindChannel(ctx context.Context, userID uint64, req BindRequest) (BindingDTO, error) {
channelID, err := strconv.ParseUint(strings.TrimSpace(req.ChannelID), 10, 64)
if err != nil || channelID == 0 {
return BindingDTO{}, errChannelIDRequired
}
code := pkgmg.NormalizeCode(req.Code)
if code == "" {
return BindingDTO{}, errCodeInvalid
}
pairing, err := repository.GetPairingCode(ctx, code)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return BindingDTO{}, errCodeInvalid
}
return BindingDTO{}, err
}
if !pairing.ExpiresAt.After(time.Now()) {
return BindingDTO{}, errCodeInvalid
}
if pairing.ChannelID != channelID {
return BindingDTO{}, errChannelMismatch
}
ch, err := repository.GetMessageChannel(ctx, channelID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return BindingDTO{}, errCodeInvalid
}
return BindingDTO{}, err
}
if !ch.Enabled {
return BindingDTO{}, errChannelDisabled
}
existing, err := repository.GetBindingByChannelPlatform(ctx, channelID, pairing.PlatformUserID)
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return BindingDTO{}, err
}
if err == nil && existing != nil {
if existing.UserID != userID {
return BindingDTO{}, errPlatformAlreadyBound
}
_ = repository.DeletePairingCode(ctx, pairing.Code)
return toBindingDTO(existing, ch), nil
}
row := &model.MessageBinding{
UserID: userID,
ChannelID: channelID,
PlatformUserID: pairing.PlatformUserID,
}
if err := repository.CreateMessageBinding(ctx, row); err != nil {
return BindingDTO{}, err
}
if err := repository.DeletePairingCode(ctx, pairing.Code); err != nil {
return BindingDTO{}, err
}
return toBindingDTO(row, ch), nil
}
// PublicChannelDTO is an enabled channel a user can bind to.
type PublicChannelDTO struct {
ID uint64 `json:"id,string"`
Name string `json:"name"`
Type string `json:"type"`
}
func listEnabledPublicChannels(ctx context.Context) ([]PublicChannelDTO, error) {
rows, err := repository.ListEnabledMessageChannels(ctx)
if err != nil {
return nil, err
}
out := make([]PublicChannelDTO, 0, len(rows))
for _, row := range rows {
out = append(out, PublicChannelDTO{ID: row.ID, Name: row.Name, Type: row.Type})
}
return out, nil
}
func listUserBindings(ctx context.Context, userID uint64) ([]BindingDTO, error) {
rows, err := repository.ListBindingsByUser(ctx, userID)
if err != nil {
return nil, err
}
out := make([]BindingDTO, 0, len(rows))
for i := range rows {
ch, err := repository.GetMessageChannel(ctx, rows[i].ChannelID)
if err != nil {
out = append(out, toBindingDTO(&rows[i], nil))
continue
}
out = append(out, toBindingDTO(&rows[i], ch))
}
return out, nil
}
func unbindChannel(ctx context.Context, userID, bindingID uint64) error {
row, err := repository.GetMessageBinding(ctx, bindingID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errBindingNotFound
}
return err
}
if row.UserID != userID {
return errBindingForbidden
}
return repository.DeleteMessageBinding(ctx, bindingID)
}
func toBindingDTO(row *model.MessageBinding, ch *model.MessageChannel) BindingDTO {
dto := BindingDTO{
ID: row.ID,
UserID: row.UserID,
ChannelID: row.ChannelID,
PlatformUserID: row.PlatformUserID,
CreatedAt: row.CreatedAt,
}
if ch != nil {
dto.ChannelName = ch.Name
dto.ChannelType = ch.Type
}
return dto
}
@@ -0,0 +1,106 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"errors"
"fmt"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"gorm.io/gorm"
)
func seedChannel(t *testing.T, ctx context.Context) *model.MessageChannel {
t.Helper()
ch := &model.MessageChannel{
Name: "tg",
Type: model.MessageChannelTypeTelegram,
OwnerScope: model.MessageOwnerScopeSystem,
Enabled: true,
}
if err := repository.CreateMessageChannel(ctx, ch); err != nil {
t.Fatalf("CreateMessageChannel() error = %v", err)
}
return ch
}
func TestBind_ExpiredCode(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
ch := seedChannel(t, ctx)
if _, err := repository.UpsertPairingCode(ctx, ch.ID, "u1", "ABCD2345", time.Now().Add(-time.Minute)); err != nil {
t.Fatalf("UpsertPairingCode() error = %v", err)
}
_, err := bindChannel(ctx, 1, BindRequest{ChannelID: fmt.Sprint(ch.ID), Code: "ABCD-2345"})
if err == nil {
t.Fatal("bindChannel() error = nil, want expired code")
}
}
func TestBind_HappyPathDeletesCode(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
ch := seedChannel(t, ctx)
if _, err := repository.UpsertPairingCode(ctx, ch.ID, "plat-1", "ABCD2345", time.Now().Add(15*time.Minute)); err != nil {
t.Fatalf("UpsertPairingCode() error = %v", err)
}
dto, err := bindChannel(ctx, 42, BindRequest{ChannelID: fmt.Sprint(ch.ID), Code: "abcd-2345"})
if err != nil {
t.Fatalf("bindChannel() error = %v", err)
}
if dto.PlatformUserID != "plat-1" || dto.ChannelID != ch.ID || dto.UserID != 42 {
t.Fatalf("bindChannel() dto = %+v", dto)
}
_, err = repository.GetPairingCode(ctx, "ABCD2345")
if !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("GetPairingCode() error = %v, want not found", err)
}
}
func TestBind_ConflictOtherUser(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
ch := seedChannel(t, ctx)
if err := repository.CreateMessageBinding(ctx, &model.MessageBinding{
UserID: 7, ChannelID: ch.ID, PlatformUserID: "plat-1",
}); err != nil {
t.Fatalf("CreateMessageBinding() error = %v", err)
}
if _, err := repository.UpsertPairingCode(ctx, ch.ID, "plat-1", "ABCD2345", time.Now().Add(15*time.Minute)); err != nil {
t.Fatalf("UpsertPairingCode() error = %v", err)
}
_, err := bindChannel(ctx, 42, BindRequest{ChannelID: fmt.Sprint(ch.ID), Code: "ABCD2345"})
if !errors.Is(err, errPlatformAlreadyBound) {
t.Fatalf("bindChannel() error = %v, want %v", err, errPlatformAlreadyBound)
}
}
func TestUnbind_OnlyOwnBinding(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
ch := seedChannel(t, ctx)
b := &model.MessageBinding{UserID: 7, ChannelID: ch.ID, PlatformUserID: "plat-1"}
if err := repository.CreateMessageBinding(ctx, b); err != nil {
t.Fatalf("CreateMessageBinding() error = %v", err)
}
if err := unbindChannel(ctx, 42, b.ID); err == nil {
t.Fatal("unbindChannel() error = nil, want forbidden")
}
if err := unbindChannel(ctx, 7, b.ID); err != nil {
t.Fatalf("unbindChannel() own binding error = %v", err)
}
_, err := repository.GetMessageBinding(ctx, b.ID)
if !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("GetMessageBinding() error = %v, want deleted", err)
}
}
+22
View File
@@ -0,0 +1,22 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package message_gateway provides user bind/unbind APIs and credential helpers.
package message_gateway
import (
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/gin-gonic/gin"
)
// RegisterUserRoutes mounts login-required bind/unbind APIs.
func RegisterUserRoutes(apiV1Router *gin.RouterGroup) {
g := apiV1Router.Group("/message-gateway")
g.Use(oauth.LoginRequired())
{
g.GET("/channels", ListChannels)
g.GET("/bindings", ListBindings)
g.POST("/bindings", BindBinding)
g.DELETE("/bindings/:id", UnbindBinding)
}
}
+14
View File
@@ -0,0 +1,14 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package v1
import (
appgw "github.com/Rain-kl/Wavelet/internal/apps/message_gateway"
"github.com/gin-gonic/gin"
)
// RegisterMessageGatewayUserRoutes mounts user bind/unbind APIs.
func RegisterMessageGatewayUserRoutes(apiV1Router *gin.RouterGroup) {
appgw.RegisterUserRoutes(apiV1Router)
}
+3
View File
@@ -16,6 +16,9 @@ func RegisterV1Routes(apiV1Router *gin.RouterGroup, apiGroup *gin.RouterGroup) {
// 2. Admin routes
RegisterAdminRoutes(apiV1Router)
// 3. Message gateway user bind/unbind
RegisterMessageGatewayUserRoutes(apiV1Router)
// 3. Product domain routes: RegisterXxxRoutes(apiV1Router) — see skill new-api
// 4. Scaffold sample only (optional demo under /api/v1/custom)
RegisterCustomRoutes(apiV1Router)