package filesystem import ( "archive/zip" "fmt" "io" "log" "net/http" "os" "path/filepath" "strings" "sync" "time" "u-desk/internal/common" "u-desk/internal/service" ) // drawio viewer 运行时下载缓存方案: // 首次预览 .drawio 时,从自建下载源拉取 drawio webapp zip → 解压到 ~/.u-desk/drawio-viewer/ // → 文件服务器 /drawio-viewer/ 从缓存目录托管静态资源 → 之后离线可用。 // 升级 drawio 时:改下面的版本号 + URL + SHA256。 const ( drawioViewerVersion = "30.3.8" // drawio 版本,升级时改这里 drawioViewerZipURL = "https://c.1216.top/u-desk/drawio-viewer-30.3.8.war" // 自建下载源(draw.war 解压即 webapp) drawioViewerSHA256 = "e12a3e190adda7c1486572bd28c1af9fffbe932df2efead3c9aa8347862505da" // draw.war v30.3.8 SHA256 drawioViewerDirName = "drawio-viewer" versionFileName = "VERSION" ) // 并发保护:多个 .drawio 同时打开时只下载一次,其余等待 var ( drawioEnsureMu sync.Mutex drawioEnsuring bool ) // DrawioViewerVersion 返回当前配置的 viewer 版本(供 binding 返回前端) func DrawioViewerVersion() string { return drawioViewerVersion } // drawioCacheDir 缓存目录 ~/.u-desk/drawio-viewer/ func drawioCacheDir() string { return filepath.Join(common.GetUserDataDir(), drawioViewerDirName) } // isDrawioViewerReady 缓存目录存在且 VERSION 文件匹配当前版本 func isDrawioViewerReady() bool { data, err := os.ReadFile(filepath.Join(drawioCacheDir(), versionFileName)) if err != nil { return false } return strings.TrimSpace(string(data)) == drawioViewerVersion } // EnsureDrawioViewer 确保本地已缓存 drawio viewer。未就绪则下载 zip + 解压。 // progress 回调推送下载进度(可为 nil)。并发安全:同时调用只下载一次,其余等待。 func EnsureDrawioViewer(progress service.DownloadProgress) error { if isDrawioViewerReady() { return nil } // 并发保护:已有线程在下载 → 轮询等待就绪 drawioEnsureMu.Lock() if drawioEnsuring { drawioEnsureMu.Unlock() for i := 0; i < 3600; i++ { // 最多等待 30 分钟(与下载超时对齐) if isDrawioViewerReady() { return nil } time.Sleep(500 * time.Millisecond) } return fmt.Errorf("等待 drawio viewer 下载超时") } drawioEnsuring = true drawioEnsureMu.Unlock() defer func() { drawioEnsureMu.Lock() drawioEnsuring = false drawioEnsureMu.Unlock() }() cacheDir := drawioCacheDir() // 1. 下载 zip(复用 service.DownloadUpdate,自带断点续传 + 哈希) res, err := service.DownloadUpdate(drawioViewerZipURL, progress) if err != nil { return fmt.Errorf("下载 drawio viewer 失败: %w", err) } // 2. SHA256 校验(若常量非空)。DownloadUpdate 成功时必返回 SHA256Hash if drawioViewerSHA256 != "" { if !strings.EqualFold(res.SHA256Hash, drawioViewerSHA256) { _ = os.Remove(res.FilePath) // 删除损坏文件,下次重试 return fmt.Errorf("drawio viewer 哈希校验失败: 期望 %s, 实际 %s", drawioViewerSHA256, res.SHA256Hash) } } // 3. 解压到临时目录,成功后原子替换缓存目录(避免半解压状态) tmpDir := cacheDir + ".tmp-" + drawioViewerVersion _ = os.RemoveAll(tmpDir) if err := unzipToDir(res.FilePath, tmpDir); err != nil { _ = os.RemoveAll(tmpDir) return fmt.Errorf("解压 drawio viewer 失败: %w", err) } // 4. 写 VERSION 文件 if err := os.WriteFile(filepath.Join(tmpDir, versionFileName), []byte(drawioViewerVersion), 0644); err != nil { _ = os.RemoveAll(tmpDir) return fmt.Errorf("写入 VERSION 失败: %w", err) } // 5. 原子替换缓存目录 _ = os.RemoveAll(cacheDir) if err := os.Rename(tmpDir, cacheDir); err != nil { _ = os.RemoveAll(tmpDir) return fmt.Errorf("移动 drawio viewer 到缓存目录失败: %w", err) } // 安装成功后删除下载的 war 包,避免长期占用磁盘(失败路径保留以支持断点续传) _ = os.Remove(res.FilePath) log.Printf("[DrawioViewer] 缓存就绪: %s (版本 %s)", cacheDir, drawioViewerVersion) return nil } // handleDrawioViewerRequest /drawio-viewer/* → 从缓存目录读静态资源(绕过文件类型白名单) // viewer 资源含 .woff2/.stencil 等白名单外扩展名,故独立路由,不走 handleLocalFileRequest。 func handleDrawioViewerRequest(w http.ResponseWriter, r *http.Request) { if writeCORSHeaders(w, r) || ensureGetMethod(w, r) { return } if !isDrawioViewerReady() { http.Error(w, "drawio viewer not ready, please call EnsureDrawioViewer first", http.StatusServiceUnavailable) return } // 剥离 /drawio-viewer/ 前缀 subPath := strings.TrimPrefix(r.URL.Path, "/drawio-viewer/") if subPath == "" { subPath = "index.html" } // 路径清理 + 防遍历:限制在缓存目录内 cleanPath, ok := sanitizeSubPath(subPath) if !ok { http.Error(w, "Forbidden", http.StatusForbidden) return } if cleanPath == "." || cleanPath == "" { cleanPath = "index.html" } absPath, ok := joinWithinBase(drawioCacheDir(), cleanPath) if !ok { http.Error(w, "Forbidden", http.StatusForbidden) return } data, err := os.ReadFile(absPath) if err != nil { http.NotFound(w, r) return } w.Header().Set("Content-Type", getDrawioViewerMIME(filepath.Ext(cleanPath))) if cleanPath == "index.html" { w.Header().Set("Cache-Control", "no-cache") // 入口短缓存,便于升级 } else { w.Header().Set("Cache-Control", "public, max-age=86400") } w.Write(data) } // sanitizeSubPath 将子路径(URL 或 zip 条目名)清理为正斜杠相对路径,含 .. 视为可疑,返回 false func sanitizeSubPath(sub string) (string, bool) { cleaned := strings.TrimPrefix(filepath.ToSlash(filepath.Clean("/"+sub)), "/") if strings.Contains(cleaned, "..") { return "", false } return cleaned, true } // joinWithinBase 将相对路径拼进 base 并确认结果仍在 base 内(防符号链接绕过),越界返回 false func joinWithinBase(base, sub string) (string, bool) { absPath := filepath.Join(base, filepath.FromSlash(sub)) if rel, err := filepath.Rel(base, absPath); err != nil || strings.HasPrefix(rel, "..") { return "", false } return absPath, true } // getDrawioViewerMIME 返回 viewer 静态资源的 MIME func getDrawioViewerMIME(ext string) string { switch strings.ToLower(ext) { case ".html", ".htm": return "text/html; charset=utf-8" case ".js", ".mjs": return "application/javascript; charset=utf-8" case ".css": return "text/css; charset=utf-8" case ".json": return "application/json" case ".svg": return "image/svg+xml" case ".png": return "image/png" case ".gif": return "image/gif" case ".jpg", ".jpeg": return "image/jpeg" case ".woff": return "font/woff" case ".woff2": return "font/woff2" case ".ttf": return "font/ttf" case ".xml": return "application/xml" default: return "application/octet-stream" } } // ==================== 辅助函数 ==================== // unzipToDir 解压整个 zip 到目标目录(覆盖已存在文件) func unzipToDir(zipPath, destDir string) error { r, err := zip.OpenReader(zipPath) if err != nil { return err } defer r.Close() if err := os.MkdirAll(destDir, 0755); err != nil { return err } for _, f := range r.File { if err := extractZipFile(f, destDir); err != nil { return err } } return nil } // extractZipFile 解压单个 zip 条目到目标目录(带路径遍历防护) func extractZipFile(f *zip.File, destDir string) error { name, ok := sanitizeSubPath(f.Name) if !ok { return fmt.Errorf("可疑路径: %s", f.Name) } destPath, ok := joinWithinBase(destDir, name) if !ok { return fmt.Errorf("逃逸路径: %s", f.Name) } if f.FileInfo().IsDir() { return os.MkdirAll(destPath, 0755) } if err := os.MkdirAll(filepath.Dir(destPath), 0755); err != nil { return err } rc, err := f.Open() if err != nil { return err } defer rc.Close() out, err := os.OpenFile(destPath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0644) if err != nil { return err } defer out.Close() _, err = io.Copy(out, rc) return err }