mirror of
https://github.com/khairul169/vaulterm.git
synced 2026-09-15 17:03:30 +07:00
feat: add auth, hosts, & keychains ownership
This commit is contained in:
+14
-10
@@ -5,10 +5,8 @@ import (
|
||||
"github.com/gofiber/fiber/v2/middleware/cors"
|
||||
"github.com/joho/godotenv"
|
||||
"rul.sh/vaulterm/app/auth"
|
||||
"rul.sh/vaulterm/app/hosts"
|
||||
"rul.sh/vaulterm/app/keychains"
|
||||
"rul.sh/vaulterm/app/ws"
|
||||
"rul.sh/vaulterm/db"
|
||||
"rul.sh/vaulterm/middleware"
|
||||
)
|
||||
|
||||
func NewApp() *fiber.App {
|
||||
@@ -18,20 +16,26 @@ func NewApp() *fiber.App {
|
||||
|
||||
// Create fiber app
|
||||
app := fiber.New(fiber.Config{ErrorHandler: ErrorHandler})
|
||||
|
||||
// Middlewares
|
||||
app.Use(cors.New())
|
||||
|
||||
// Init app routes
|
||||
auth.Router(app)
|
||||
hosts.Router(app)
|
||||
keychains.Router(app)
|
||||
ws.Router(app)
|
||||
// Server info
|
||||
app.Get("/server", func(c *fiber.Ctx) error {
|
||||
return c.JSON(&fiber.Map{
|
||||
"name": "Vaulterm",
|
||||
"version": "0.0.1",
|
||||
})
|
||||
})
|
||||
|
||||
// Health check
|
||||
app.Get("/health-check", func(c *fiber.Ctx) error {
|
||||
return c.SendString("OK")
|
||||
})
|
||||
|
||||
app.Use(middleware.Auth)
|
||||
auth.Router(app)
|
||||
|
||||
app.Use(middleware.Protected())
|
||||
InitRouter(app)
|
||||
|
||||
return app
|
||||
}
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
|
||||
type Auth struct{ db *gorm.DB }
|
||||
|
||||
func NewAuthRepository() *Auth {
|
||||
func NewRepository() *Auth {
|
||||
return &Auth{db: db.Get()}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,10 +1,9 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"rul.sh/vaulterm/lib"
|
||||
"rul.sh/vaulterm/middleware"
|
||||
"rul.sh/vaulterm/utils"
|
||||
)
|
||||
|
||||
@@ -12,12 +11,12 @@ func Router(app *fiber.App) {
|
||||
router := app.Group("/auth")
|
||||
|
||||
router.Post("/login", login)
|
||||
router.Get("/user", getUser)
|
||||
router.Post("/logout", logout)
|
||||
router.Get("/user", middleware.Protected(), getUser)
|
||||
router.Post("/logout", middleware.Protected(), logout)
|
||||
}
|
||||
|
||||
func login(c *fiber.Ctx) error {
|
||||
repo := NewAuthRepository()
|
||||
repo := NewRepository()
|
||||
|
||||
var body LoginSchema
|
||||
if err := c.BodyParser(&body); err != nil {
|
||||
@@ -54,32 +53,15 @@ func login(c *fiber.Ctx) error {
|
||||
}
|
||||
|
||||
func getUser(c *fiber.Ctx) error {
|
||||
auth := c.Get("Authorization")
|
||||
var sessionId string
|
||||
|
||||
if auth != "" {
|
||||
sessionId = strings.Split(auth, " ")[1]
|
||||
}
|
||||
|
||||
repo := NewAuthRepository()
|
||||
session, err := repo.GetSession(sessionId)
|
||||
if err != nil {
|
||||
return utils.ResponseError(c, err, 500)
|
||||
}
|
||||
|
||||
return c.JSON(session)
|
||||
user := utils.GetUser(c)
|
||||
return c.JSON(user)
|
||||
}
|
||||
|
||||
func logout(c *fiber.Ctx) error {
|
||||
auth := c.Get("Authorization")
|
||||
force := c.Query("force")
|
||||
var sessionId string
|
||||
sessionId := c.Locals("sessionId").(string)
|
||||
|
||||
if auth != "" {
|
||||
sessionId = strings.Split(auth, " ")[1]
|
||||
}
|
||||
|
||||
repo := NewAuthRepository()
|
||||
repo := NewRepository()
|
||||
err := repo.RemoveUserSession(sessionId, force == "true")
|
||||
|
||||
if err != nil {
|
||||
|
||||
@@ -4,25 +4,36 @@ import (
|
||||
"gorm.io/gorm"
|
||||
"rul.sh/vaulterm/db"
|
||||
"rul.sh/vaulterm/models"
|
||||
"rul.sh/vaulterm/utils"
|
||||
)
|
||||
|
||||
type Hosts struct{ db *gorm.DB }
|
||||
type Hosts struct {
|
||||
db *gorm.DB
|
||||
User *utils.UserContext
|
||||
}
|
||||
|
||||
func NewRepository() *Hosts {
|
||||
return &Hosts{db: db.Get()}
|
||||
func NewRepository(r *Hosts) *Hosts {
|
||||
if r == nil {
|
||||
r = &Hosts{}
|
||||
}
|
||||
r.db = db.Get()
|
||||
return r
|
||||
}
|
||||
|
||||
func (r *Hosts) GetAll() ([]*models.Host, error) {
|
||||
query := r.ACL(r.db.Order("id DESC"))
|
||||
|
||||
var rows []*models.Host
|
||||
ret := r.db.Order("id DESC").Find(&rows)
|
||||
ret := query.Find(&rows)
|
||||
|
||||
return rows, ret.Error
|
||||
}
|
||||
|
||||
func (r *Hosts) Get(id string) (*models.HostDecrypted, error) {
|
||||
var host models.Host
|
||||
ret := r.db.Joins("Key").Joins("AltKey").Where("hosts.id = ?", id).First(&host)
|
||||
query := r.ACL(r.db)
|
||||
|
||||
var host models.Host
|
||||
ret := query.Joins("Key").Joins("AltKey").Where("hosts.id = ?", id).First(&host)
|
||||
if ret.Error != nil {
|
||||
return nil, ret.Error
|
||||
}
|
||||
@@ -37,12 +48,13 @@ func (r *Hosts) Get(id string) (*models.HostDecrypted, error) {
|
||||
|
||||
func (r *Hosts) Exists(id string) (bool, error) {
|
||||
var count int64
|
||||
ret := r.db.Model(&models.Host{}).Where("id = ?", id).Count(&count)
|
||||
ret := r.ACL(r.db.Model(&models.Host{}).Where("id = ?", id)).Count(&count)
|
||||
return count > 0, ret.Error
|
||||
}
|
||||
|
||||
func (r *Hosts) Delete(id string) error {
|
||||
return r.db.Delete(&models.Host{Model: models.Model{ID: id}}).Error
|
||||
query := r.ACL(r.db)
|
||||
return query.Delete(&models.Host{Model: models.Model{ID: id}}).Error
|
||||
}
|
||||
|
||||
func (r *Hosts) Create(item *models.Host) error {
|
||||
@@ -50,5 +62,15 @@ func (r *Hosts) Create(item *models.Host) error {
|
||||
}
|
||||
|
||||
func (r *Hosts) Update(id string, item *models.Host) error {
|
||||
return r.db.Where("id = ?", id).Updates(item).Error
|
||||
query := r.ACL(r.db.Where("id = ?", id))
|
||||
|
||||
return query.Updates(item).Error
|
||||
}
|
||||
|
||||
func (r *Hosts) ACL(query *gorm.DB) *gorm.DB {
|
||||
if r.User.IsAdmin {
|
||||
return query
|
||||
}
|
||||
|
||||
return query.Where("hosts.owner_id = ?", r.User.ID)
|
||||
}
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
"rul.sh/vaulterm/utils"
|
||||
)
|
||||
|
||||
func Router(app *fiber.App) {
|
||||
func Router(app fiber.Router) {
|
||||
router := app.Group("/hosts")
|
||||
|
||||
router.Get("/", getAll)
|
||||
@@ -19,7 +19,9 @@ func Router(app *fiber.App) {
|
||||
}
|
||||
|
||||
func getAll(c *fiber.Ctx) error {
|
||||
repo := NewRepository()
|
||||
user := utils.GetUser(c)
|
||||
repo := NewRepository(&Hosts{User: user})
|
||||
|
||||
rows, err := repo.GetAll()
|
||||
if err != nil {
|
||||
return utils.ResponseError(c, err, 500)
|
||||
@@ -36,8 +38,11 @@ func create(c *fiber.Ctx) error {
|
||||
return utils.ResponseError(c, err, 500)
|
||||
}
|
||||
|
||||
repo := NewRepository()
|
||||
user := utils.GetUser(c)
|
||||
repo := NewRepository(&Hosts{User: user})
|
||||
|
||||
item := &models.Host{
|
||||
OwnerID: user.ID,
|
||||
Type: body.Type,
|
||||
Label: body.Label,
|
||||
Host: body.Host,
|
||||
@@ -48,7 +53,7 @@ func create(c *fiber.Ctx) error {
|
||||
AltKeyID: body.AltKeyID,
|
||||
}
|
||||
|
||||
osName, err := tryConnect(item)
|
||||
osName, err := tryConnect(c, item)
|
||||
if err != nil {
|
||||
return utils.ResponseError(c, fmt.Errorf("cannot connect to the host: %s", err), 500)
|
||||
}
|
||||
@@ -67,7 +72,8 @@ func update(c *fiber.Ctx) error {
|
||||
return utils.ResponseError(c, err, 500)
|
||||
}
|
||||
|
||||
repo := NewRepository()
|
||||
user := utils.GetUser(c)
|
||||
repo := NewRepository(&Hosts{User: user})
|
||||
|
||||
id := c.Params("id")
|
||||
exist, _ := repo.Exists(id)
|
||||
@@ -87,7 +93,7 @@ func update(c *fiber.Ctx) error {
|
||||
AltKeyID: body.AltKeyID,
|
||||
}
|
||||
|
||||
osName, err := tryConnect(item)
|
||||
osName, err := tryConnect(c, item)
|
||||
if err != nil {
|
||||
return utils.ResponseError(c, fmt.Errorf("cannot connect to the host: %s", err), 500)
|
||||
}
|
||||
@@ -101,7 +107,8 @@ func update(c *fiber.Ctx) error {
|
||||
}
|
||||
|
||||
func delete(c *fiber.Ctx) error {
|
||||
repo := NewRepository()
|
||||
user := utils.GetUser(c)
|
||||
repo := NewRepository(&Hosts{User: user})
|
||||
|
||||
id := c.Params("id")
|
||||
exist, _ := repo.Exists(id)
|
||||
|
||||
@@ -3,13 +3,16 @@ package hosts
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"rul.sh/vaulterm/app/keychains"
|
||||
"rul.sh/vaulterm/lib"
|
||||
"rul.sh/vaulterm/models"
|
||||
"rul.sh/vaulterm/utils"
|
||||
)
|
||||
|
||||
func tryConnect(host *models.Host) (string, error) {
|
||||
keyRepo := keychains.NewRepository()
|
||||
func tryConnect(c *fiber.Ctx, host *models.Host) (string, error) {
|
||||
user := utils.GetUser(c)
|
||||
keyRepo := keychains.NewRepository(&keychains.Keychains{User: user})
|
||||
|
||||
var key map[string]interface{}
|
||||
var altKey map[string]interface{}
|
||||
|
||||
@@ -4,18 +4,27 @@ import (
|
||||
"gorm.io/gorm"
|
||||
"rul.sh/vaulterm/db"
|
||||
"rul.sh/vaulterm/models"
|
||||
"rul.sh/vaulterm/utils"
|
||||
)
|
||||
|
||||
type Keychains struct{ db *gorm.DB }
|
||||
type Keychains struct {
|
||||
db *gorm.DB
|
||||
User *utils.UserContext
|
||||
}
|
||||
|
||||
func NewRepository() *Keychains {
|
||||
return &Keychains{db: db.Get()}
|
||||
func NewRepository(r *Keychains) *Keychains {
|
||||
if r == nil {
|
||||
r = &Keychains{}
|
||||
}
|
||||
r.db = db.Get()
|
||||
return r
|
||||
}
|
||||
|
||||
func (r *Keychains) GetAll() ([]*models.Keychain, error) {
|
||||
var rows []*models.Keychain
|
||||
ret := r.db.Order("created_at DESC").Find(&rows)
|
||||
query := r.ACL(r.db.Order("created_at DESC"))
|
||||
|
||||
ret := query.Find(&rows)
|
||||
return rows, ret.Error
|
||||
}
|
||||
|
||||
@@ -25,7 +34,9 @@ func (r *Keychains) Create(item *models.Keychain) error {
|
||||
|
||||
func (r *Keychains) Get(id string) (*models.Keychain, error) {
|
||||
var keychain models.Keychain
|
||||
if err := r.db.Where("id = ?", id).First(&keychain).Error; err != nil {
|
||||
query := r.ACL(r.db.Where("id = ?", id))
|
||||
|
||||
if err := query.First(&keychain).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -34,7 +45,8 @@ func (r *Keychains) Get(id string) (*models.Keychain, error) {
|
||||
|
||||
func (r *Keychains) Exists(id string) (bool, error) {
|
||||
var count int64
|
||||
ret := r.db.Model(&models.Keychain{}).Where("id = ?", id).Count(&count)
|
||||
query := r.ACL(r.db.Model(&models.Keychain{}).Where("id = ?", id))
|
||||
ret := query.Count(&count)
|
||||
return count > 0, ret.Error
|
||||
}
|
||||
|
||||
@@ -58,5 +70,14 @@ func (r *Keychains) GetDecrypted(id string) (*KeychainDecrypted, error) {
|
||||
}
|
||||
|
||||
func (r *Keychains) Update(id string, item *models.Keychain) error {
|
||||
return r.db.Where("id = ?", id).Updates(item).Error
|
||||
query := r.ACL(r.db.Where("id = ?", id))
|
||||
return query.Updates(item).Error
|
||||
}
|
||||
|
||||
func (r *Keychains) ACL(query *gorm.DB) *gorm.DB {
|
||||
if r.User.IsAdmin {
|
||||
return query
|
||||
}
|
||||
|
||||
return query.Where("keychains.owner_id = ?", r.User.ID)
|
||||
}
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
"rul.sh/vaulterm/utils"
|
||||
)
|
||||
|
||||
func Router(app *fiber.App) {
|
||||
func Router(app fiber.Router) {
|
||||
router := app.Group("/keychains")
|
||||
|
||||
router.Get("/", getAll)
|
||||
@@ -25,7 +25,9 @@ type GetAllResult struct {
|
||||
func getAll(c *fiber.Ctx) error {
|
||||
withData := c.Query("withData")
|
||||
|
||||
repo := NewRepository()
|
||||
user := utils.GetUser(c)
|
||||
repo := NewRepository(&Keychains{User: user})
|
||||
|
||||
rows, err := repo.GetAll()
|
||||
if err != nil {
|
||||
return utils.ResponseError(c, err, 500)
|
||||
@@ -62,11 +64,13 @@ func create(c *fiber.Ctx) error {
|
||||
return utils.ResponseError(c, err, 500)
|
||||
}
|
||||
|
||||
repo := NewRepository()
|
||||
user := utils.GetUser(c)
|
||||
repo := NewRepository(&Keychains{User: user})
|
||||
|
||||
item := &models.Keychain{
|
||||
Type: body.Type,
|
||||
Label: body.Label,
|
||||
OwnerID: user.ID,
|
||||
Type: body.Type,
|
||||
Label: body.Label,
|
||||
}
|
||||
|
||||
if err := item.EncryptData(body.Data); err != nil {
|
||||
@@ -86,7 +90,9 @@ func update(c *fiber.Ctx) error {
|
||||
return utils.ResponseError(c, err, 500)
|
||||
}
|
||||
|
||||
repo := NewRepository()
|
||||
user := utils.GetUser(c)
|
||||
repo := NewRepository(&Keychains{User: user})
|
||||
|
||||
id := c.Params("id")
|
||||
|
||||
exist, _ := repo.Exists(id)
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"rul.sh/vaulterm/app/hosts"
|
||||
"rul.sh/vaulterm/app/keychains"
|
||||
"rul.sh/vaulterm/app/ws"
|
||||
)
|
||||
|
||||
func InitRouter(app *fiber.App) {
|
||||
// App route list
|
||||
routes := []Router{
|
||||
hosts.Router,
|
||||
keychains.Router,
|
||||
ws.Router,
|
||||
}
|
||||
|
||||
for _, route := range routes {
|
||||
route(app)
|
||||
}
|
||||
}
|
||||
|
||||
type Router func(app fiber.Router)
|
||||
@@ -5,7 +5,7 @@ import (
|
||||
"github.com/gofiber/fiber/v2"
|
||||
)
|
||||
|
||||
func Router(app *fiber.App) {
|
||||
func Router(app fiber.Router) {
|
||||
router := app.Group("/ws")
|
||||
|
||||
router.Use(func(c *fiber.Ctx) error {
|
||||
|
||||
@@ -13,7 +13,8 @@ import (
|
||||
func HandleTerm(c *websocket.Conn) {
|
||||
hostId := c.Query("hostId")
|
||||
|
||||
hostRepo := hosts.NewRepository()
|
||||
user := utils.GetUserWs(c)
|
||||
hostRepo := hosts.NewRepository(&hosts.Hosts{User: user})
|
||||
data, err := hostRepo.Get(hostId)
|
||||
|
||||
if data == nil {
|
||||
|
||||
Reference in New Issue
Block a user