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) } }