package workers import ( "archive/zip" "fmt" "io" "os" "path/filepath" "strings" "sync" "time" "git.elijahkuntz.com/Elijah/drive/config" "github.com/google/uuid" ) type TaskStatus string const ( TaskPending TaskStatus = "pending" TaskRunning TaskStatus = "running" TaskCompleted TaskStatus = "completed" TaskFailed TaskStatus = "failed" ) type BackgroundTask struct { ID string `json:"id"` Type string `json:"type"` // "zip" or "unzip" Status TaskStatus `json:"status"` Progress int `json:"progress"` // 0-100 Message string `json:"message"` Error string `json:"error,omitempty"` CreatedAt time.Time `json:"created_at"` } type TaskManager struct { Tasks map[string]*BackgroundTask mu sync.RWMutex cfg *config.Config } var Tasks *TaskManager func InitTaskManager(cfg *config.Config) { Tasks = &TaskManager{ Tasks: make(map[string]*BackgroundTask), cfg: cfg, } } func (m *TaskManager) CreateTask(taskType string, message string) *BackgroundTask { id := uuid.New().String() task := &BackgroundTask{ ID: id, Type: taskType, Status: TaskPending, Progress: 0, Message: message, CreatedAt: time.Now(), } m.mu.Lock() m.Tasks[id] = task m.mu.Unlock() return task } func (m *TaskManager) UpdateTask(id string, status TaskStatus, progress int, msg string) { m.mu.Lock() defer m.mu.Unlock() if task, exists := m.Tasks[id]; exists { task.Status = status task.Progress = progress if msg != "" { task.Message = msg } } } func (m *TaskManager) FailTask(id string, err error) { m.mu.Lock() defer m.mu.Unlock() if task, exists := m.Tasks[id]; exists { task.Status = TaskFailed task.Error = err.Error() } } func (m *TaskManager) GetTasks() []BackgroundTask { m.mu.RLock() defer m.mu.RUnlock() // Convert map to sorted slice (newest first) var list []BackgroundTask for _, t := range m.Tasks { list = append(list, *t) } return list } func (m *TaskManager) CleanOldTasks() { m.mu.Lock() defer m.mu.Unlock() for id, t := range m.Tasks { if (t.Status == TaskCompleted || t.Status == TaskFailed) && time.Since(t.CreatedAt) > 24*time.Hour { delete(m.Tasks, id) } } } // --- ZIP / UNZIP Logic --- func ZipFolderAsync(taskID string, sourceFull, destFull string) { Tasks.UpdateTask(taskID, TaskRunning, 0, "Counting files...") // Count total files for progress calculation var totalFiles int filepath.Walk(sourceFull, func(_ string, info os.FileInfo, err error) error { if err == nil && !info.IsDir() { totalFiles++ } return nil }) if totalFiles == 0 { Tasks.FailTask(taskID, fmt.Errorf("folder is empty")) return } Tasks.UpdateTask(taskID, TaskRunning, 5, "Creating archive...") // Create zip file zipFile, err := os.Create(destFull) if err != nil { Tasks.FailTask(taskID, err) return } defer zipFile.Close() archive := zip.NewWriter(zipFile) defer archive.Close() var processedFiles int filepath.Walk(sourceFull, func(path string, info os.FileInfo, err error) error { if err != nil { return err } // Calculate relative path inside the zip relPath, err := filepath.Rel(filepath.Dir(sourceFull), path) if err != nil { return err } if info.IsDir() { // Create directory entry _, err = archive.Create(relPath + "/") return err } // Add file file, err := os.Open(path) if err != nil { return err } defer file.Close() w, err := archive.Create(relPath) if err != nil { return err } _, err = io.Copy(w, file) if err != nil { return err } processedFiles++ progress := 5 + int(float64(processedFiles)/float64(totalFiles)*90) Tasks.UpdateTask(taskID, TaskRunning, progress, fmt.Sprintf("Compressing %s...", info.Name())) return nil }) Tasks.UpdateTask(taskID, TaskCompleted, 100, "Archive created successfully") } func UnzipAsync(taskID string, sourceFull, destFolder string) { Tasks.UpdateTask(taskID, TaskRunning, 0, "Opening archive...") r, err := zip.OpenReader(sourceFull) if err != nil { Tasks.FailTask(taskID, err) return } defer r.Close() totalFiles := len(r.File) if totalFiles == 0 { Tasks.FailTask(taskID, fmt.Errorf("archive is empty")) return } Tasks.UpdateTask(taskID, TaskRunning, 5, "Extracting files...") for i, f := range r.File { // Prevent Zip Slip vulnerability cleanPath := filepath.Clean(f.Name) if strings.HasPrefix(cleanPath, "..") { continue // Skip malicious paths } fpath := filepath.Join(destFolder, cleanPath) if f.FileInfo().IsDir() { os.MkdirAll(fpath, os.ModePerm) continue } if err = os.MkdirAll(filepath.Dir(fpath), os.ModePerm); err != nil { Tasks.FailTask(taskID, err) return } outFile, err := os.OpenFile(fpath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, f.Mode()) if err != nil { Tasks.FailTask(taskID, err) return } rc, err := f.Open() if err != nil { outFile.Close() Tasks.FailTask(taskID, err) return } _, err = io.Copy(outFile, rc) outFile.Close() rc.Close() if err != nil { Tasks.FailTask(taskID, err) return } progress := 5 + int(float64(i+1)/float64(totalFiles)*90) Tasks.UpdateTask(taskID, TaskRunning, progress, fmt.Sprintf("Extracted %s", filepath.Base(fpath))) } Tasks.UpdateTask(taskID, TaskCompleted, 100, "Extracted successfully") }