package middleware import ( "crypto/sha256" "encoding/hex" "encoding/json" "net/http" "net/http/httptest" "testing" "github.com/gin-gonic/gin" "github.com/gochat/gochat/internal/auth" "github.com/gochat/gochat/internal/model" "github.com/stretchr/testify/require" "gorm.io/driver/sqlite" "gorm.io/gorm" ) func TestConnectorServiceTokenIsLimitedToApplicationAllowlistAndAccountGrant(t *testing.T) { gin.SetMode(gin.TestMode) db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{}) require.NoError(t, err) require.NoError(t, db.AutoMigrate(&model.PlatformApp{}, &model.AccessToken{}, &model.Permissible{})) active := true app := &model.PlatformApp{Name: "SWT", Type: "integration", Status: "active", Active: &active, Config: json.RawMessage(`{"connector":"shangwutong"}`)} require.NoError(t, db.Create(app).Error) token := "gochat_pa_connector_allowlist" digest := sha256.Sum256([]byte(token)) require.NoError(t, db.Create(&model.AccessToken{ OwnerType: model.AccessTokenOwnerTypePlatformApp, OwnerID: app.ID, Token: hex.EncodeToString(digest[:]), TokenPrefix: token[:8], Name: "connector", }).Error) require.NoError(t, db.Create(&model.Permissible{ PlatformAppID: app.ID, PermissibleType: model.PermissibleTypeAccount, PermissibleID: 1, }).Error) router := gin.New() api := router.Group("/api/v1") api.Use(AuthMiddlewareWithConnectorAllowlist(auth.NewJWTService(makeJWTConfig()), db)) accounts := api.Group("/accounts") accounts.Use(AccountScope()) accounts.POST("/:account_id/conversations/:conversation_id/messages", func(c *gin.Context) { c.Status(http.StatusOK) }) accounts.GET("/:account_id/conversations/:conversation_id/messages", func(c *gin.Context) { c.Status(http.StatusOK) }) accounts.POST("/:account_id/inboxes", func(c *gin.Context) { c.Status(http.StatusOK) }) request := func(method, path string) int { response := httptest.NewRecorder() req := httptest.NewRequest(method, path, nil) req.Header.Set("Authorization", "Bearer "+token) router.ServeHTTP(response, req) return response.Code } require.Equal(t, http.StatusOK, request(http.MethodPost, "/api/v1/accounts/1/conversations/2/messages")) require.Equal(t, http.StatusForbidden, request(http.MethodPost, "/api/v1/accounts/2/conversations/2/messages")) require.Equal(t, http.StatusUnauthorized, request(http.MethodGet, "/api/v1/accounts/1/conversations/2/messages")) require.Equal(t, http.StatusUnauthorized, request(http.MethodPost, "/api/v1/accounts/1/inboxes")) }