feat(backend): user + library admin API, atomic sanitized upload
This commit is contained in:
@@ -0,0 +1,167 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"io"
|
||||
"log"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"booklib/internal/bookfile"
|
||||
"booklib/internal/store"
|
||||
)
|
||||
|
||||
// resolveLibRoot: root_path 必须绝对且落在 BooksDir 内(spec §7 前缀校验)
|
||||
func (a *api) libRoot(c *gin.Context, lib store.Library) (string, bool) {
|
||||
root := filepath.Clean(lib.RootPath)
|
||||
books := filepath.Clean(a.cfg.BooksDir)
|
||||
if !filepath.IsAbs(root) || (root != books && !strings.HasPrefix(root, books+string(os.PathSeparator))) {
|
||||
err(c, http.StatusForbidden, "forbidden", "library root outside books dir")
|
||||
return "", false
|
||||
}
|
||||
return root, true
|
||||
}
|
||||
|
||||
func (a *api) listLibraries(c *gin.Context) {
|
||||
libs, e := a.st.ListLibraries(c)
|
||||
if e != nil {
|
||||
log.Printf("db: %v", e)
|
||||
err(c, http.StatusInternalServerError, "internal", "db error")
|
||||
return
|
||||
}
|
||||
out := make([]gin.H, 0, len(libs))
|
||||
for _, l := range libs {
|
||||
out = append(out, gin.H{"id": l.ID, "name": l.Name, "root_path": l.RootPath,
|
||||
"created_at": l.CreatedAt.Format(time.RFC3339)})
|
||||
}
|
||||
c.JSON(http.StatusOK, out)
|
||||
}
|
||||
|
||||
func (a *api) createLibrary(c *gin.Context) {
|
||||
var req struct {
|
||||
Name string `json:"name"`
|
||||
RootPath string `json:"root_path"`
|
||||
}
|
||||
if c.ShouldBindJSON(&req) != nil || req.Name == "" || req.RootPath == "" {
|
||||
err(c, http.StatusBadRequest, "bad_request", "name and root_path required")
|
||||
return
|
||||
}
|
||||
if !filepath.IsAbs(req.RootPath) {
|
||||
err(c, http.StatusBadRequest, "bad_request", "root_path must be absolute")
|
||||
return
|
||||
}
|
||||
id, e := a.st.CreateLibrary(c, req.Name, filepath.Clean(req.RootPath))
|
||||
if e != nil {
|
||||
if isUnique(e) {
|
||||
err(c, http.StatusConflict, "exists", "root_path taken")
|
||||
return
|
||||
}
|
||||
err(c, http.StatusBadRequest, "bad_request", "invalid input")
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusCreated, gin.H{"id": id, "name": req.Name, "root_path": filepath.Clean(req.RootPath)})
|
||||
}
|
||||
|
||||
func (a *api) getLibrary(c *gin.Context) (store.Library, bool) {
|
||||
id, e := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if e != nil {
|
||||
err(c, http.StatusBadRequest, "bad_request", "bad id")
|
||||
return store.Library{}, false
|
||||
}
|
||||
lib, e := a.st.GetLibrary(c, id)
|
||||
if e != nil {
|
||||
err(c, http.StatusNotFound, "not_found", "no such library")
|
||||
return store.Library{}, false
|
||||
}
|
||||
return lib, true
|
||||
}
|
||||
|
||||
func (a *api) scanLibrary(c *gin.Context) {
|
||||
if _, ok := a.getLibrary(c); !ok {
|
||||
return
|
||||
}
|
||||
// Task 9 接线: go a.sc.ScanLibraryByID(...)
|
||||
err(c, http.StatusNotImplemented, "not_ready", "scanner not wired yet")
|
||||
}
|
||||
|
||||
func (a *api) upload(c *gin.Context) {
|
||||
lib, ok := a.getLibrary(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
root, ok := a.libRoot(c, lib)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, a.cfg.UploadMaxMB<<20)
|
||||
fh, e := c.FormFile("file")
|
||||
if e != nil {
|
||||
err(c, http.StatusBadRequest, "bad_request", "multipart field 'file' required")
|
||||
return
|
||||
}
|
||||
name := bookfile.SafeName(fh.Filename)
|
||||
if bookfile.FormatFromExt(name) == "" {
|
||||
err(c, http.StatusBadRequest, "bad_format", "extension must be cbz/pdf/epub/txt/md")
|
||||
return
|
||||
}
|
||||
dst, e := a.uniquePath(root, name)
|
||||
if e != nil {
|
||||
err(c, http.StatusForbidden, "forbidden", e.Error())
|
||||
return
|
||||
}
|
||||
src, e := fh.Open()
|
||||
if e != nil {
|
||||
err(c, http.StatusInternalServerError, "internal", "open upload")
|
||||
return
|
||||
}
|
||||
defer src.Close()
|
||||
tmp := dst + ".upload-" + strconv.FormatInt(time.Now().UnixNano(), 36)
|
||||
out, e := os.OpenFile(tmp, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o644)
|
||||
if e != nil {
|
||||
err(c, http.StatusInternalServerError, "internal", "create tmp")
|
||||
return
|
||||
}
|
||||
if _, e := io.Copy(out, src); e != nil {
|
||||
out.Close()
|
||||
os.Remove(tmp)
|
||||
err(c, http.StatusRequestEntityTooLarge, "too_large", "upload failed")
|
||||
return
|
||||
}
|
||||
out.Close()
|
||||
if e := os.Rename(tmp, dst); e != nil { // 原子落盘,scanner 自动收编
|
||||
os.Remove(tmp)
|
||||
err(c, http.StatusInternalServerError, "internal", "rename")
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusAccepted, gin.H{"accepted": true, "path": strings.TrimPrefix(dst, root+string(os.PathSeparator))})
|
||||
}
|
||||
|
||||
// uniquePath 清洗后的 name 必须仍在 root 内;重名加 " (n)" 后缀
|
||||
func (a *api) uniquePath(root, name string) (string, error) {
|
||||
ext := filepath.Ext(name)
|
||||
base := strings.TrimSuffix(name, ext)
|
||||
for i := 0; ; i++ {
|
||||
cand := base + ext
|
||||
if i > 0 {
|
||||
cand = base + " (" + strconv.Itoa(i) + ")" + ext
|
||||
}
|
||||
p := filepath.Join(root, cand)
|
||||
if filepath.Clean(p) != filepath.Join(root, filepath.Clean(cand)) ||
|
||||
!strings.HasPrefix(filepath.Clean(p), root+string(os.PathSeparator)) {
|
||||
return "", os.ErrInvalid
|
||||
}
|
||||
if _, e := os.Stat(p); os.IsNotExist(e) {
|
||||
return p, nil
|
||||
} else if e != nil {
|
||||
return "", e
|
||||
}
|
||||
if i > 999 {
|
||||
return "", os.ErrExist
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"mime/multipart"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestLibraryCreateListUpload(t *testing.T) {
|
||||
_, h := setupAPI(t)
|
||||
tok := adminToken(t, h)
|
||||
root := t.TempDir()
|
||||
w := do(h, "POST", "/api/libraries", tok, map[string]string{"name": "comics", "root_path": root})
|
||||
if w.Code != 201 {
|
||||
t.Fatalf("create lib %d %s", w.Code, w.Body)
|
||||
}
|
||||
var lib map[string]any
|
||||
json.Unmarshal(w.Body.Bytes(), &lib)
|
||||
libID := itoa(lib["id"])
|
||||
w = do(h, "GET", "/api/libraries", tok, nil)
|
||||
if !strings.Contains(w.Body.String(), `"comics"`) {
|
||||
t.Fatalf("list: %s", w.Body)
|
||||
}
|
||||
// 相对路径 root 必须 400(前缀校验的根)
|
||||
w = do(h, "POST", "/api/libraries", tok, map[string]string{"name": "x", "root_path": "relative/path"})
|
||||
if w.Code != 400 {
|
||||
t.Fatalf("relative root want 400 got %d", w.Code)
|
||||
}
|
||||
// 上传:白名单 + 防穿越 + 原子落盘
|
||||
body, mw := uploadBody("my 01.cbz", []byte("zipbytes"))
|
||||
req := httptest.NewRequest("POST", "/api/libraries/"+libID+"/upload", body)
|
||||
req.Header.Set("Content-Type", mw.FormDataContentType())
|
||||
req.Header.Set("Authorization", "Bearer "+tok)
|
||||
ww := httptest.NewRecorder()
|
||||
h.ServeHTTP(ww, req)
|
||||
if ww.Code != 202 {
|
||||
t.Fatalf("upload %d %s", ww.Code, ww.Body)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(root, "my 01.cbz")); err != nil {
|
||||
t.Fatal("uploaded file missing:", err)
|
||||
}
|
||||
body, mw = uploadBody("../../evil.cbz", []byte("x"))
|
||||
req = httptest.NewRequest("POST", "/api/libraries/"+libID+"/upload", body)
|
||||
req.Header.Set("Content-Type", mw.FormDataContentType())
|
||||
req.Header.Set("Authorization", "Bearer "+tok)
|
||||
ww = httptest.NewRecorder()
|
||||
h.ServeHTTP(ww, req)
|
||||
if ww.Code != 202 { // 名字被清洗成 evil.cbz,落在 root 内
|
||||
t.Fatalf("sanitize upload %d", ww.Code)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(root, "evil.cbz")); err != nil {
|
||||
t.Fatal("evil upload not sanitized")
|
||||
}
|
||||
body, mw = uploadBody("virus.exe", []byte("x"))
|
||||
req = httptest.NewRequest("POST", "/api/libraries/"+libID+"/upload", body)
|
||||
req.Header.Set("Content-Type", mw.FormDataContentType())
|
||||
req.Header.Set("Authorization", "Bearer "+tok)
|
||||
ww = httptest.NewRecorder()
|
||||
h.ServeHTTP(ww, req)
|
||||
if ww.Code != 400 {
|
||||
t.Fatalf("bad ext want 400 got %d", ww.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func uploadBody(filename string, content []byte) (*bytes.Buffer, *multipart.Writer) {
|
||||
buf := &bytes.Buffer{}
|
||||
mw := multipart.NewWriter(buf)
|
||||
fw, _ := mw.CreateFormFile("file", filename)
|
||||
fw.Write(content)
|
||||
mw.Close()
|
||||
return buf, mw
|
||||
}
|
||||
@@ -21,7 +21,16 @@ func NewRouter(cfg *config.Config, st *store.Store, rdb *redispkg.R) *gin.Engine
|
||||
|
||||
p := g.Group("", a.authMw())
|
||||
p.GET("/auth/me", a.me)
|
||||
// ponytail: 501 垫片,Task 5 换成 a.createUser;无此路由则 member 403 测不到
|
||||
p.POST("/users", a.adminOnly(), func(c *gin.Context) { c.Status(http.StatusNotImplemented) })
|
||||
|
||||
users := p.Group("/users", a.adminOnly())
|
||||
users.GET("", a.listUsers)
|
||||
users.POST("", a.createUser)
|
||||
users.DELETE("/:id", a.deleteUser)
|
||||
|
||||
libs := p.Group("/libraries")
|
||||
libs.GET("", a.listLibraries)
|
||||
libs.POST("", a.adminOnly(), a.createLibrary)
|
||||
libs.POST("/:id/scan", a.adminOnly(), a.scanLibrary)
|
||||
libs.POST("/:id/upload", a.adminOnly(), a.upload)
|
||||
return r
|
||||
}
|
||||
|
||||
@@ -0,0 +1,99 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"log"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgconn"
|
||||
|
||||
"booklib/internal/auth"
|
||||
)
|
||||
|
||||
func isUnique(e error) bool {
|
||||
var pgErr *pgconn.PgError
|
||||
return errors.As(e, &pgErr) && pgErr.Code == "23505"
|
||||
}
|
||||
|
||||
func (a *api) listUsers(c *gin.Context) {
|
||||
users, e := a.st.ListUsers(c)
|
||||
if e != nil {
|
||||
log.Printf("db: %v", e)
|
||||
err(c, http.StatusInternalServerError, "internal", "db error")
|
||||
return
|
||||
}
|
||||
out := make([]gin.H, 0, len(users))
|
||||
for _, u := range users {
|
||||
out = append(out, gin.H{"id": u.ID, "username": u.Username, "role": u.Role,
|
||||
"created_at": u.CreatedAt.Format(time.RFC3339)})
|
||||
}
|
||||
c.JSON(http.StatusOK, out)
|
||||
}
|
||||
|
||||
func (a *api) createUser(c *gin.Context) {
|
||||
var req struct{ Username, Password, Role string }
|
||||
if c.ShouldBindJSON(&req) != nil {
|
||||
err(c, http.StatusBadRequest, "bad_request", "json body required")
|
||||
return
|
||||
}
|
||||
if req.Role != "admin" && req.Role != "member" {
|
||||
err(c, http.StatusBadRequest, "bad_request", "role must be admin|member")
|
||||
return
|
||||
}
|
||||
if len(req.Password) < 8 {
|
||||
err(c, http.StatusBadRequest, "bad_request", "password too short (min 8)")
|
||||
return
|
||||
}
|
||||
h, e := auth.HashPassword(req.Password)
|
||||
if e != nil {
|
||||
err(c, http.StatusInternalServerError, "internal", "hash")
|
||||
return
|
||||
}
|
||||
id, e := a.st.CreateUser(c, req.Username, h, req.Role)
|
||||
if e != nil {
|
||||
if isUnique(e) {
|
||||
err(c, http.StatusConflict, "exists", "username taken")
|
||||
return
|
||||
}
|
||||
err(c, http.StatusBadRequest, "bad_request", "invalid input")
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusCreated, gin.H{"id": id, "username": req.Username, "role": req.Role})
|
||||
}
|
||||
|
||||
func (a *api) deleteUser(c *gin.Context) {
|
||||
id, e := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if e != nil {
|
||||
err(c, http.StatusBadRequest, "bad_request", "bad id")
|
||||
return
|
||||
}
|
||||
if id == uid(c) {
|
||||
err(c, http.StatusBadRequest, "bad_request", "cannot delete yourself")
|
||||
return
|
||||
}
|
||||
target, e := a.st.GetUserByID(c, id)
|
||||
if e != nil {
|
||||
if errors.Is(e, pgx.ErrNoRows) {
|
||||
err(c, http.StatusNotFound, "not_found", "no such user")
|
||||
return
|
||||
}
|
||||
err(c, http.StatusInternalServerError, "internal", "db error")
|
||||
return
|
||||
}
|
||||
if target.Role == "admin" {
|
||||
n, _ := a.st.CountAdmins(c) // 防删光最后一个 admin
|
||||
if n <= 1 {
|
||||
err(c, http.StatusBadRequest, "bad_request", "cannot delete the last admin")
|
||||
return
|
||||
}
|
||||
}
|
||||
if e := a.st.DeleteUser(c, id); e != nil {
|
||||
err(c, http.StatusInternalServerError, "internal", "db error")
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func adminToken(t *testing.T, h http.Handler) string {
|
||||
t.Helper()
|
||||
w := do(h, "POST", "/api/auth/login", "", map[string]string{"username": "alice", "password": "pw12345"})
|
||||
var v struct{ Token string }
|
||||
json.Unmarshal(w.Body.Bytes(), &v)
|
||||
return v.Token
|
||||
}
|
||||
|
||||
func TestUserCRUD(t *testing.T) {
|
||||
_, h := setupAPI(t)
|
||||
tok := adminToken(t, h)
|
||||
w := do(h, "POST", "/api/users", tok, map[string]string{"username": "carol", "password": "pw123456", "role": "member"})
|
||||
if w.Code != 201 {
|
||||
t.Fatalf("create %d %s", w.Code, w.Body)
|
||||
}
|
||||
w = do(h, "GET", "/api/users", tok, nil)
|
||||
var users []map[string]any
|
||||
json.Unmarshal(w.Body.Bytes(), &users)
|
||||
if len(users) != 3 {
|
||||
t.Fatalf("list want 3 got %d: %s", len(users), w.Body)
|
||||
}
|
||||
carolID := findID(users, "carol")
|
||||
// 重名 → 409
|
||||
w = do(h, "POST", "/api/users", tok, map[string]string{"username": "carol", "password": "pw123456", "role": "member"})
|
||||
if w.Code != 409 {
|
||||
t.Fatalf("dup want 409 got %d", w.Code)
|
||||
}
|
||||
// 弱密码 → 400
|
||||
w = do(h, "POST", "/api/users", tok, map[string]string{"username": "dave", "password": "1", "role": "member"})
|
||||
if w.Code != 400 {
|
||||
t.Fatalf("weak want 400 got %d", w.Code)
|
||||
}
|
||||
// 不能删自己:先 me 拿 id
|
||||
w = do(h, "GET", "/api/auth/me", tok, nil)
|
||||
var me map[string]any
|
||||
json.Unmarshal(w.Body.Bytes(), &me)
|
||||
w = do(h, "DELETE", "/api/users/"+itoa(me["id"]), tok, nil)
|
||||
if w.Code != 400 {
|
||||
t.Fatalf("self-delete want 400 got %d %s", w.Code, w.Body)
|
||||
}
|
||||
// 删 carol → 204,再删 → 404
|
||||
w = do(h, "DELETE", "/api/users/"+itoa(carolID), tok, nil)
|
||||
if w.Code != 204 {
|
||||
t.Fatalf("delete %d %s", w.Code, w.Body)
|
||||
}
|
||||
w = do(h, "DELETE", "/api/users/"+itoa(carolID), tok, nil)
|
||||
if w.Code != 404 {
|
||||
t.Fatalf("redelete want 404 got %d", w.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBadRoleRejected(t *testing.T) {
|
||||
_, h := setupAPI(t)
|
||||
tok := adminToken(t, h)
|
||||
w := do(h, "POST", "/api/users", tok, map[string]string{"username": "e", "password": "pw123456", "role": "god"})
|
||||
if w.Code != 400 {
|
||||
t.Fatalf("bad role want 400 got %d", w.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func findID(rows []map[string]any, name string) float64 {
|
||||
for _, r := range rows {
|
||||
if r["username"] == name {
|
||||
return r["id"].(float64)
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func itoa(v any) string { return strconv.FormatFloat(v.(float64), 'f', 0, 64) }
|
||||
Reference in New Issue
Block a user