Some checks failed
Automated Container Build / build-and-push (push) Failing after 12s
236 lines
6.6 KiB
Go
236 lines
6.6 KiB
Go
package handlers
|
|
|
|
import (
|
|
"fmt"
|
|
"os"
|
|
"path/filepath"
|
|
"time"
|
|
|
|
"git.elijahkuntz.com/Elijah/drive/config"
|
|
"git.elijahkuntz.com/Elijah/drive/database"
|
|
"github.com/gofiber/fiber/v2"
|
|
"github.com/google/uuid"
|
|
"golang.org/x/crypto/bcrypt"
|
|
)
|
|
|
|
type ShareHandler struct {
|
|
DB *database.DB
|
|
Config *config.Config
|
|
}
|
|
|
|
type ShareRequest struct {
|
|
Path string `json:"path"`
|
|
Password string `json:"password,omitempty"`
|
|
ExpiresIn int `json:"expires_in,omitempty"` // Hours
|
|
MaxDownloads int `json:"max_downloads,omitempty"`
|
|
}
|
|
|
|
// CreateShare creates a new public share link.
|
|
// POST /api/share
|
|
func (h *ShareHandler) CreateShare(c *fiber.Ctx) error {
|
|
var req ShareRequest
|
|
if err := c.BodyParser(&req); err != nil {
|
|
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{"error": "invalid request"})
|
|
}
|
|
|
|
// Validate path
|
|
fullPath := filepath.Join(h.Config.StorageDir, filepath.FromSlash(req.Path))
|
|
if _, err := os.Stat(fullPath); err != nil {
|
|
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{"error": "file not found"})
|
|
}
|
|
|
|
id := uuid.New().String()
|
|
var hash *string
|
|
if req.Password != "" {
|
|
hashed, err := bcrypt.GenerateFromPassword([]byte(req.Password), 10)
|
|
if err == nil {
|
|
s := string(hashed)
|
|
hash = &s
|
|
}
|
|
}
|
|
|
|
var expiresAt *time.Time
|
|
if req.ExpiresIn > 0 {
|
|
t := time.Now().Add(time.Duration(req.ExpiresIn) * time.Hour)
|
|
expiresAt = &t
|
|
}
|
|
|
|
var maxDownloads *int
|
|
if req.MaxDownloads > 0 {
|
|
maxDownloads = &req.MaxDownloads
|
|
}
|
|
|
|
_, err := h.DB.Exec(`
|
|
INSERT INTO share_links (id, file_path, password_hash, expires_at, max_downloads)
|
|
VALUES (?, ?, ?, ?, ?)
|
|
`, id, req.Path, hash, expiresAt, maxDownloads)
|
|
|
|
if err != nil {
|
|
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "failed to create share"})
|
|
}
|
|
|
|
h.DB.AddAuditLog("share_created", fmt.Sprintf("Shared %s (ID: %s)", req.Path, id), c.IP())
|
|
|
|
return c.JSON(fiber.Map{
|
|
"id": id,
|
|
"link": fmt.Sprintf("/share/%s", id), // Frontend route
|
|
})
|
|
}
|
|
|
|
// ListShares returns all active shares.
|
|
// GET /api/share
|
|
func (h *ShareHandler) ListShares(c *fiber.Ctx) error {
|
|
rows, err := h.DB.Query(`
|
|
SELECT id, file_path, expires_at, downloads, max_downloads, created_at,
|
|
(password_hash IS NOT NULL) as has_password
|
|
FROM share_links
|
|
ORDER BY created_at DESC
|
|
`)
|
|
if err != nil {
|
|
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "db error"})
|
|
}
|
|
defer rows.Close()
|
|
|
|
var shares []map[string]interface{}
|
|
for rows.Next() {
|
|
var id, path, createdAt string
|
|
var expiresAt *string
|
|
var downloads int
|
|
var maxDownloads *int
|
|
var hasPassword bool
|
|
|
|
rows.Scan(&id, &path, &expiresAt, &downloads, &maxDownloads, &createdAt, &hasPassword)
|
|
|
|
shares = append(shares, map[string]interface{}{
|
|
"id": id,
|
|
"path": path,
|
|
"expires_at": expiresAt,
|
|
"downloads": downloads,
|
|
"max_downloads": maxDownloads,
|
|
"created_at": createdAt,
|
|
"has_password": hasPassword,
|
|
})
|
|
}
|
|
|
|
return c.JSON(fiber.Map{"shares": shares})
|
|
}
|
|
|
|
// RevokeShare removes a share.
|
|
// DELETE /api/share/:id
|
|
func (h *ShareHandler) RevokeShare(c *fiber.Ctx) error {
|
|
id := c.Params("id")
|
|
_, err := h.DB.Exec(`DELETE FROM share_links WHERE id = ?`, id)
|
|
if err != nil {
|
|
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "failed to revoke"})
|
|
}
|
|
|
|
h.DB.AddAuditLog("share_revoked", fmt.Sprintf("Revoked share %s", id), c.IP())
|
|
return c.JSON(fiber.Map{"message": "revoked"})
|
|
}
|
|
|
|
// GetPublicShare returns share info without auth (used by the public page).
|
|
// POST /api/public/share/:id
|
|
// Body might contain {"password": "..."}
|
|
func (h *ShareHandler) GetPublicShare(c *fiber.Ctx) error {
|
|
id := c.Params("id")
|
|
var req struct {
|
|
Password string `json:"password"`
|
|
}
|
|
c.BodyParser(&req)
|
|
|
|
var path string
|
|
var hash *string
|
|
var expiresAt *time.Time
|
|
var maxDownloads *int
|
|
var downloads int
|
|
|
|
err := h.DB.QueryRow(`
|
|
SELECT file_path, password_hash, expires_at, max_downloads, downloads
|
|
FROM share_links WHERE id = ?
|
|
`, id).Scan(&path, &hash, &expiresAt, &maxDownloads, &downloads)
|
|
|
|
if err != nil {
|
|
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{"error": "share not found"})
|
|
}
|
|
|
|
// Check expiry
|
|
if expiresAt != nil && time.Now().After(*expiresAt) {
|
|
return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "link expired"})
|
|
}
|
|
|
|
// Check download limits (we only count actual downloads, not views)
|
|
if maxDownloads != nil && downloads >= *maxDownloads {
|
|
return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "download limit reached"})
|
|
}
|
|
|
|
// Check password
|
|
if hash != nil {
|
|
if req.Password == "" {
|
|
return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{"require_password": true})
|
|
}
|
|
if err := bcrypt.CompareHashAndPassword([]byte(*hash), []byte(req.Password)); err != nil {
|
|
return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{"error": "invalid password"})
|
|
}
|
|
}
|
|
|
|
// Send file info
|
|
fullPath := filepath.Join(h.Config.StorageDir, filepath.FromSlash(path))
|
|
info, err := os.Stat(fullPath)
|
|
if err != nil {
|
|
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{"error": "file no longer exists"})
|
|
}
|
|
|
|
// We return a signed download token or just allow the immediate request.
|
|
// For simplicity, we'll return a short-lived download token JWT.
|
|
// Actually, easier: the frontend calls GET /api/public/download/:id?pwd=xxx
|
|
|
|
return c.JSON(fiber.Map{
|
|
"name": filepath.Base(path),
|
|
"size": info.Size(),
|
|
"is_dir": info.IsDir(),
|
|
})
|
|
}
|
|
|
|
// DownloadPublicShare serves the file.
|
|
// GET /api/public/download/:id
|
|
func (h *ShareHandler) DownloadPublicShare(c *fiber.Ctx) error {
|
|
id := c.Params("id")
|
|
pwd := c.Query("pwd")
|
|
|
|
var path string
|
|
var hash *string
|
|
var expiresAt *time.Time
|
|
var maxDownloads *int
|
|
var downloads int
|
|
|
|
err := h.DB.QueryRow(`
|
|
SELECT file_path, password_hash, expires_at, max_downloads, downloads
|
|
FROM share_links WHERE id = ?
|
|
`, id).Scan(&path, &hash, &expiresAt, &maxDownloads, &downloads)
|
|
|
|
if err != nil {
|
|
return c.Status(fiber.StatusNotFound).SendString("not found")
|
|
}
|
|
|
|
if expiresAt != nil && time.Now().After(*expiresAt) {
|
|
return c.Status(fiber.StatusForbidden).SendString("expired")
|
|
}
|
|
|
|
if maxDownloads != nil && downloads >= *maxDownloads {
|
|
return c.Status(fiber.StatusForbidden).SendString("limit reached")
|
|
}
|
|
|
|
if hash != nil {
|
|
if err := bcrypt.CompareHashAndPassword([]byte(*hash), []byte(pwd)); err != nil {
|
|
return c.Status(fiber.StatusUnauthorized).SendString("unauthorized")
|
|
}
|
|
}
|
|
|
|
fullPath := filepath.Join(h.Config.StorageDir, filepath.FromSlash(path))
|
|
|
|
// Increment download counter
|
|
h.DB.Exec(`UPDATE share_links SET downloads = downloads + 1 WHERE id = ?`, id)
|
|
h.DB.AddAuditLog("share_downloaded", fmt.Sprintf("Downloaded share %s", id), c.IP())
|
|
|
|
return c.SendFile(fullPath)
|
|
}
|