Files
zsp-project/backend/test/company_handler_test.go
2026-06-03 20:59:39 +08:00

496 lines
15 KiB
Go

package test
import (
"bytes"
"encoding/json"
"errors"
"mime/multipart"
"net/http"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"github.com/rovina/zsp-backend/internal/handler"
"github.com/rovina/zsp-backend/internal/model"
"github.com/rovina/zsp-backend/internal/service"
)
// mockCompanyService implements service.CompanyService for testing
type mockCompanyService struct {
createFolderFn func(name, description, category string) (*model.CompanyFolder, error)
listFoldersFn func() ([]model.CompanyFolder, error)
getFolderByIDFn func(id uint) (*model.CompanyFolder, error)
updateFolderFn func(id uint, name, description, category string) error
deleteFolderFn func(id uint) error
uploadFilesFn func(folderID *uint, companyName, fileCategory, expireDate, description string, files []*multipart.FileHeader) (int, error)
listFilesFn func(folderID *uint) ([]model.CompanyFile, error)
getFileByIDFn func(id uint) (*model.CompanyFile, error)
updateFileFn func(id uint, companyName, fileCategory, expireDate, description string, folderID *uint) error
deleteFileFn func(id uint) error
}
func (m *mockCompanyService) CreateFolder(name, description, category string) (*model.CompanyFolder, error) {
return m.createFolderFn(name, description, category)
}
func (m *mockCompanyService) ListFolders() ([]model.CompanyFolder, error) {
return m.listFoldersFn()
}
func (m *mockCompanyService) GetFolderByID(id uint) (*model.CompanyFolder, error) {
return m.getFolderByIDFn(id)
}
func (m *mockCompanyService) UpdateFolder(id uint, name, description, category string) error {
return m.updateFolderFn(id, name, description, category)
}
func (m *mockCompanyService) DeleteFolder(id uint) error {
return m.deleteFolderFn(id)
}
func (m *mockCompanyService) UploadFiles(folderID *uint, companyName, fileCategory, expireDate, description string, files []*multipart.FileHeader) (int, error) {
return m.uploadFilesFn(folderID, companyName, fileCategory, expireDate, description, files)
}
func (m *mockCompanyService) ListFiles(folderID *uint) ([]model.CompanyFile, error) {
return m.listFilesFn(folderID)
}
func (m *mockCompanyService) GetFileByID(id uint) (*model.CompanyFile, error) {
return m.getFileByIDFn(id)
}
func (m *mockCompanyService) UpdateFile(id uint, companyName, fileCategory, expireDate, description string, folderID *uint) error {
return m.updateFileFn(id, companyName, fileCategory, expireDate, description, folderID)
}
func (m *mockCompanyService) DeleteFile(id uint) error {
return m.deleteFileFn(id)
}
var _ service.CompanyService = (*mockCompanyService)(nil)
func setupCompanyRouter(h *handler.CompanyHandler) *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
api := r.Group("/api")
{
company := api.Group("/company")
{
company.GET("/folders", h.ListFolders)
company.POST("/folders", h.CreateFolder)
company.PUT("/folders/:id", h.UpdateFolder)
company.DELETE("/folders/:id", h.DeleteFolder)
company.GET("/files", h.ListFiles)
company.POST("/files/upload", h.UploadFiles)
company.PUT("/files/:id", h.UpdateFile)
company.DELETE("/files/:id", h.DeleteFile)
}
}
return r
}
// ========== Folder Tests ==========
func TestCompanyListFolders_Success(t *testing.T) {
mock := &mockCompanyService{
listFoldersFn: func() ([]model.CompanyFolder, error) {
return []model.CompanyFolder{
{ID: 1, Name: "营业执照文件夹", Category: "license"},
{ID: 2, Name: "开票信息文件夹", Category: "invoice"},
}, nil
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/api/company/folders", nil)
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", w.Code)
}
var resp map[string]interface{}
json.Unmarshal(w.Body.Bytes(), &resp)
if resp["success"] != true {
t.Errorf("expected success true, got %v", resp["success"])
}
}
func TestCompanyListFolders_Error(t *testing.T) {
mock := &mockCompanyService{
listFoldersFn: func() ([]model.CompanyFolder, error) {
return nil, errors.New("database error")
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/api/company/folders", nil)
router.ServeHTTP(w, req)
if w.Code != http.StatusInternalServerError {
t.Errorf("expected status 500, got %d", w.Code)
}
}
func TestCompanyCreateFolder_Success(t *testing.T) {
mock := &mockCompanyService{
createFolderFn: func(name, description, category string) (*model.CompanyFolder, error) {
return &model.CompanyFolder{ID: 1, Name: name, Description: description, Category: category}, nil
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
body := `{"name":"测试文件夹","description":"测试描述","category":"license"}`
w := httptest.NewRecorder()
req, _ := http.NewRequest("POST", "/api/company/folders", bytes.NewReader([]byte(body)))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", w.Code)
}
var resp map[string]interface{}
json.Unmarshal(w.Body.Bytes(), &resp)
if resp["success"] != true {
t.Errorf("expected success true, got %v", resp["success"])
}
}
func TestCompanyCreateFolder_BadRequest(t *testing.T) {
mock := &mockCompanyService{
createFolderFn: func(name, description, category string) (*model.CompanyFolder, error) {
return nil, nil
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
w := httptest.NewRecorder()
req, _ := http.NewRequest("POST", "/api/company/folders", bytes.NewReader([]byte("{}")))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400 for missing required field, got %d", w.Code)
}
}
func TestCompanyCreateFolder_Error(t *testing.T) {
mock := &mockCompanyService{
createFolderFn: func(name, description, category string) (*model.CompanyFolder, error) {
return nil, errors.New("folder already exists")
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
body := `{"name":"测试文件夹","description":"测试描述","category":"license"}`
w := httptest.NewRecorder()
req, _ := http.NewRequest("POST", "/api/company/folders", bytes.NewReader([]byte(body)))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)
if w.Code != http.StatusInternalServerError {
t.Errorf("expected status 500, got %d", w.Code)
}
}
func TestCompanyUpdateFolder_Success(t *testing.T) {
mock := &mockCompanyService{
updateFolderFn: func(id uint, name, description, category string) error {
return nil
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
body := `{"name":"更新文件夹","description":"更新描述","category":"invoice"}`
w := httptest.NewRecorder()
req, _ := http.NewRequest("PUT", "/api/company/folders/1", bytes.NewReader([]byte(body)))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", w.Code)
}
}
func TestCompanyUpdateFolder_InvalidID(t *testing.T) {
mock := &mockCompanyService{
updateFolderFn: func(id uint, name, description, category string) error {
return nil
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
body := `{"name":"更新文件夹","description":"更新描述","category":"invoice"}`
w := httptest.NewRecorder()
req, _ := http.NewRequest("PUT", "/api/company/folders/invalid", bytes.NewReader([]byte(body)))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400 for invalid id, got %d", w.Code)
}
}
func TestCompanyUpdateFolder_Error(t *testing.T) {
mock := &mockCompanyService{
updateFolderFn: func(id uint, name, description, category string) error {
return errors.New("folder not found")
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
body := `{"name":"更新文件夹","description":"更新描述","category":"invoice"}`
w := httptest.NewRecorder()
req, _ := http.NewRequest("PUT", "/api/company/folders/999", bytes.NewReader([]byte(body)))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)
if w.Code != http.StatusInternalServerError {
t.Errorf("expected status 500, got %d", w.Code)
}
}
func TestCompanyDeleteFolder_Success(t *testing.T) {
mock := &mockCompanyService{
deleteFolderFn: func(id uint) error {
return nil
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
w := httptest.NewRecorder()
req, _ := http.NewRequest("DELETE", "/api/company/folders/1", nil)
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", w.Code)
}
}
func TestCompanyDeleteFolder_InvalidID(t *testing.T) {
mock := &mockCompanyService{
deleteFolderFn: func(id uint) error {
return nil
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
w := httptest.NewRecorder()
req, _ := http.NewRequest("DELETE", "/api/company/folders/invalid", nil)
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400 for invalid id, got %d", w.Code)
}
}
func TestCompanyDeleteFolder_Error(t *testing.T) {
mock := &mockCompanyService{
deleteFolderFn: func(id uint) error {
return errors.New("folder not found")
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
w := httptest.NewRecorder()
req, _ := http.NewRequest("DELETE", "/api/company/folders/999", nil)
router.ServeHTTP(w, req)
if w.Code != http.StatusInternalServerError {
t.Errorf("expected status 500, got %d", w.Code)
}
}
// ========== File Tests ==========
func TestCompanyListFiles_Success(t *testing.T) {
mock := &mockCompanyService{
listFilesFn: func(folderID *uint) ([]model.CompanyFile, error) {
return []model.CompanyFile{
{ID: 1, FileName: "营业执照.pdf", CompanyName: "公司A", FileCategory: "license"},
{ID: 2, FileName: "开票信息.pdf", CompanyName: "公司A", FileCategory: "invoice"},
}, nil
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/api/company/files", nil)
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", w.Code)
}
var resp map[string]interface{}
json.Unmarshal(w.Body.Bytes(), &resp)
if resp["success"] != true {
t.Errorf("expected success true, got %v", resp["success"])
}
}
func TestCompanyListFiles_WithFolderID(t *testing.T) {
mock := &mockCompanyService{
listFilesFn: func(folderID *uint) ([]model.CompanyFile, error) {
if folderID != nil && *folderID != 1 {
t.Errorf("expected folderID 1, got %v", folderID)
}
return []model.CompanyFile{
{ID: 1, FileName: "营业执照.pdf", FolderID: 1},
}, nil
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/api/company/files?folderId=1", nil)
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", w.Code)
}
}
func TestCompanyListFiles_Error(t *testing.T) {
mock := &mockCompanyService{
listFilesFn: func(folderID *uint) ([]model.CompanyFile, error) {
return nil, errors.New("database error")
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/api/company/files", nil)
router.ServeHTTP(w, req)
if w.Code != http.StatusInternalServerError {
t.Errorf("expected status 500, got %d", w.Code)
}
}
func TestCompanyUpdateFile_Success(t *testing.T) {
mock := &mockCompanyService{
updateFileFn: func(id uint, companyName, fileCategory, expireDate, description string, folderID *uint) error {
return nil
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
body := `{"companyName":"公司B","fileCategory":"license","expireDate":"2025-12-31","description":"更新描述"}`
w := httptest.NewRecorder()
req, _ := http.NewRequest("PUT", "/api/company/files/1", bytes.NewReader([]byte(body)))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", w.Code)
}
}
func TestCompanyUpdateFile_InvalidID(t *testing.T) {
mock := &mockCompanyService{
updateFileFn: func(id uint, companyName, fileCategory, expireDate, description string, folderID *uint) error {
return nil
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
body := `{"companyName":"公司B","fileCategory":"license"}`
w := httptest.NewRecorder()
req, _ := http.NewRequest("PUT", "/api/company/files/invalid", bytes.NewReader([]byte(body)))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400 for invalid id, got %d", w.Code)
}
}
func TestCompanyUpdateFile_Error(t *testing.T) {
mock := &mockCompanyService{
updateFileFn: func(id uint, companyName, fileCategory, expireDate, description string, folderID *uint) error {
return errors.New("file not found")
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
body := `{"companyName":"公司B","fileCategory":"license"}`
w := httptest.NewRecorder()
req, _ := http.NewRequest("PUT", "/api/company/files/999", bytes.NewReader([]byte(body)))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)
if w.Code != http.StatusInternalServerError {
t.Errorf("expected status 500, got %d", w.Code)
}
}
func TestCompanyDeleteFile_Success(t *testing.T) {
mock := &mockCompanyService{
deleteFileFn: func(id uint) error {
return nil
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
w := httptest.NewRecorder()
req, _ := http.NewRequest("DELETE", "/api/company/files/1", nil)
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", w.Code)
}
}
func TestCompanyDeleteFile_InvalidID(t *testing.T) {
mock := &mockCompanyService{
deleteFileFn: func(id uint) error {
return nil
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
w := httptest.NewRecorder()
req, _ := http.NewRequest("DELETE", "/api/company/files/invalid", nil)
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400 for invalid id, got %d", w.Code)
}
}
func TestCompanyDeleteFile_Error(t *testing.T) {
mock := &mockCompanyService{
deleteFileFn: func(id uint) error {
return errors.New("file not found")
},
}
h := handler.NewCompanyHandler(mock)
router := setupCompanyRouter(h)
w := httptest.NewRecorder()
req, _ := http.NewRequest("DELETE", "/api/company/files/999", nil)
router.ServeHTTP(w, req)
if w.Code != http.StatusInternalServerError {
t.Errorf("expected status 500, got %d", w.Code)
}
}