package npm_test import ( "archive/tar" "bytes" "compress/gzip" "context" "crypto/sha512" "encoding/base64" "encoding/json" "fmt" "net/http" "net/http/httptest" "os" "path/filepath" "strings" "testing" "git.fsdpf.net/go/jscriptx/esm/npm" ) // 假注册表:把包元数据和 tarball 都放内存里,测试不碰网络。 type fakeRegistry struct { *httptest.Server // 包名 -> 版本 -> 依赖 pkgs map[string]map[string]map[string]string // 故意破坏某个包的校验和 corrupt map[string]bool } func newRegistry(t *testing.T) *fakeRegistry { t.Helper() r := &fakeRegistry{ pkgs: map[string]map[string]map[string]string{}, corrupt: map[string]bool{}, } mux := http.NewServeMux() mux.HandleFunc("/", func(w http.ResponseWriter, req *http.Request) { path := strings.TrimPrefix(req.URL.Path, "/") // //-/-.tgz 取 tarball if i := strings.Index(path, "/-/"); i >= 0 { name := path[:i] file := strings.TrimSuffix(path[i+3:], ".tgz") ver := file[strings.LastIndexByte(file, '-')+1:] body := r.tarball(t, name, ver) if r.corrupt[name] { body = append(body, 0) // 改一个字节,校验和就对不上了 } w.Write(body) return } // / 取元数据 vers, ok := r.pkgs[path] if !ok { http.NotFound(w, req) return } doc := map[string]any{"dist-tags": map[string]string{"latest": highest(vers)}} out := map[string]any{} for v, deps := range vers { out[v] = map[string]any{ "dependencies": deps, "dist": map[string]string{ "tarball": r.URL + "/" + path + "/-/" + path + "-" + v + ".tgz", "integrity": integrityOf(r.tarball(t, path, v)), }, } } doc["versions"] = out json.NewEncoder(w).Encode(doc) }) r.Server = httptest.NewServer(mux) t.Cleanup(r.Close) return r } // add 往假注册表里放一个版本。 func (r *fakeRegistry) add(name, version string, deps map[string]string) { if r.pkgs[name] == nil { r.pkgs[name] = map[string]map[string]string{} } r.pkgs[name][version] = deps } // tarball 现造一个 npm 格式的 tarball:外面套一层 package/。 func (r *fakeRegistry) tarball(t *testing.T, name, version string) []byte { t.Helper() pkg, _ := json.Marshal(map[string]any{ "name": name, "version": version, "main": "index.js", "dependencies": r.pkgs[name][version], }) files := map[string]string{ "package/package.json": string(pkg), "package/index.js": fmt.Sprintf("export const who = %q\n", name+"@"+version), } var buf bytes.Buffer zw := gzip.NewWriter(&buf) tw := tar.NewWriter(zw) for _, n := range []string{"package/package.json", "package/index.js"} { body := files[n] tw.WriteHeader(&tar.Header{Name: n, Mode: 0o644, Size: int64(len(body)), Typeflag: tar.TypeReg}) tw.Write([]byte(body)) } tw.Close() zw.Close() return buf.Bytes() } func integrityOf(b []byte) string { sum := sha512.Sum512(b) return "sha512-" + base64.StdEncoding.EncodeToString(sum[:]) } func highest(vers map[string]map[string]string) string { best := "" for v := range vers { if v > best { // 测试里的版本号都是同位数的,字符串比较够用 best = v } } return best } func client(r *fakeRegistry) *npm.Client { return &npm.Client{Registry: r.URL, HTTP: r.Client()} } // 依赖树要被完整拉下来,平铺到 node_modules。 func TestFetch_把依赖树平铺下来(t *testing.T) { r := newRegistry(t) r.add("main", "1.0.0", map[string]string{"mid": "^2.0.0"}) r.add("mid", "2.3.0", map[string]string{"leaf": "~1.1.0"}) r.add("leaf", "1.1.5", nil) r.add("leaf", "1.2.0", nil) // ~1.1.0 不该选到它 dir := t.TempDir() got, err := client(r).Fetch(context.Background(), dir, "main") if err != nil { t.Fatal(err) } want := map[string]string{"main": "1.0.0", "mid": "2.3.0", "leaf": "1.1.5"} if len(got) != len(want) { t.Fatalf("拉下来 %d 个包,该是 %d 个:%+v", len(got), len(want), got) } for _, p := range got { if want[p.Name] != p.Version { t.Errorf("%s 装的是 %s,该是 %s", p.Name, p.Version, want[p.Name]) } // 平铺:全在 node_modules 顶层 if filepath.Dir(p.Dir) != filepath.Join(dir, "node_modules") { t.Errorf("%s 没在顶层:%s", p.Name, p.Dir) } if _, err := os.Stat(filepath.Join(p.Dir, "package.json")); err != nil { t.Errorf("%s 没解开: %v", p.Name, err) } } } // 同一个包被多处依赖,只装一次。 func TestFetch_共同依赖只装一次(t *testing.T) { r := newRegistry(t) r.add("a", "1.0.0", map[string]string{"shared": "^1.0.0"}) r.add("b", "1.0.0", map[string]string{"shared": "^1.2.0"}) r.add("shared", "1.5.0", nil) got, err := client(r).Fetch(context.Background(), t.TempDir(), "a", "b") if err != nil { t.Fatal(err) } n := 0 for _, p := range got { if p.Name == "shared" { n++ } } if n != 1 { t.Errorf("shared 装了 %d 次,该只装 1 次", n) } } // 主版本冲突要报错说清楚,而不是随便挑一个装上。 func TestFetch_版本冲突要报错(t *testing.T) { r := newRegistry(t) r.add("a", "1.0.0", map[string]string{"dep": "^1.0.0"}) r.add("b", "1.0.0", map[string]string{"dep": "^2.0.0"}) r.add("dep", "1.0.0", nil) r.add("dep", "2.0.0", nil) _, err := client(r).Fetch(context.Background(), t.TempDir(), "a", "b") if err == nil { t.Fatal("该报冲突") } for _, want := range []string{"dep", "冲突", "esm.Vendor"} { if !strings.Contains(err.Error(), want) { t.Errorf("错误里该提到 %q: %v", want, err) } } } // 校验和对不上就得拒绝——中间的缓存代理、私有源镜像都可能改动内容。 func TestFetch_校验和对不上就拒绝(t *testing.T) { r := newRegistry(t) r.add("bad", "1.0.0", nil) r.corrupt["bad"] = true _, err := client(r).Fetch(context.Background(), t.TempDir(), "bad") if err == nil { t.Fatal("校验和不对时该报错") } if !strings.Contains(err.Error(), "校验和") { t.Errorf("错误该说是校验和的问题: %v", err) } } // 不写版本就取 latest。 func TestFetch_不写版本取latest(t *testing.T) { r := newRegistry(t) r.add("p", "1.0.0", nil) r.add("p", "1.4.0", nil) got, err := client(r).Fetch(context.Background(), t.TempDir(), "p") if err != nil { t.Fatal(err) } if got[0].Version != "1.4.0" { t.Errorf("装了 %s,该是 latest 1.4.0", got[0].Version) } } // 指定范围要生效。 func TestFetch_按范围挑版本(t *testing.T) { r := newRegistry(t) r.add("p", "1.0.0", nil) r.add("p", "1.4.0", nil) r.add("p", "2.0.0", nil) got, err := client(r).Fetch(context.Background(), t.TempDir(), "p@^1") if err != nil { t.Fatal(err) } if got[0].Version != "1.4.0" { t.Errorf("装了 %s,该是 1.4.0", got[0].Version) } } // 包不存在要说清楚是哪个包。 func TestFetch_包不存在(t *testing.T) { r := newRegistry(t) _, err := client(r).Fetch(context.Background(), t.TempDir(), "nope") if err == nil || !strings.Contains(err.Error(), "nope") { t.Errorf("该报错并提到包名: %v", err) } } // tarball 是外来数据,路径里带 ../ 的要挡住,不能写出目录之外。 func TestFetch_挡住跳出目录的路径(t *testing.T) { var buf bytes.Buffer zw := gzip.NewWriter(&buf) tw := tar.NewWriter(zw) body := "被写到外面去了" tw.WriteHeader(&tar.Header{ Name: "package/../../../evil.js", Mode: 0o644, Size: int64(len(body)), Typeflag: tar.TypeReg, }) tw.Write([]byte(body)) tw.Close() zw.Close() evil := buf.Bytes() srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { if strings.Contains(req.URL.Path, "/-/") { w.Write(evil) return } json.NewEncoder(w).Encode(map[string]any{ "dist-tags": map[string]string{"latest": "1.0.0"}, "versions": map[string]any{"1.0.0": map[string]any{ "dist": map[string]string{ "tarball": "http://" + req.Host + "/evil/-/evil-1.0.0.tgz", "integrity": integrityOf(evil), }}}, }) })) defer srv.Close() c := &npm.Client{Registry: srv.URL, HTTP: srv.Client()} if _, err := c.Fetch(context.Background(), t.TempDir(), "evil"); err == nil { t.Fatal("带 ../ 的路径该被挡住") } else if !strings.Contains(err.Error(), "可疑路径") { t.Errorf("错误该说明是路径问题: %v", err) } }