From 4b0b96e87e32fd67e993b8fde6e045edfbf5a3eb Mon Sep 17 00:00:00 2001 From: qmuntal Date: Fri, 4 Sep 2026 11:33:18 +0200 Subject: [PATCH 1/3] Use Copilot CLI release assets in Go bundler --- go/README.md | 4 +- go/cmd/bundler/main.go | 557 +++++++++++++++---------- go/cmd/bundler/main_test.go | 311 +++++++++++++- go/internal/embeddedcli/embeddedcli.go | 137 ++---- 4 files changed, 666 insertions(+), 343 deletions(-) diff --git a/go/README.md b/go/README.md index 801ad8556c..f329261f07 100644 --- a/go/README.md +++ b/go/README.md @@ -107,9 +107,11 @@ Follow these steps to embed the CLI: 1. Run `go get -tool github.com/github/copilot-sdk/go/cmd/bundler`. This is a one-time setup step per project. 2. Run `go tool bundler` in your build environment just before building your application. +The bundler downloads the pinned `github/copilot-cli` release archive and verifies it against that release's `SHA256SUMS.txt` before extracting any files. Set `COPILOT_CLI_DOWNLOAD_BASE_URL` to use a release mirror while packaging. + That's it! When your application calls `copilot.NewClient` without a `Connection` field (or with an empty `StdioConnection{}`), the SDK automatically installs the embedded `copilot-runtime` executable and adjacent `runtime.node` to a cache directory for managed child-process connections. -The bundler prepares the native runtime library required by the [in-process transport](#in-process-transport-experimental). It is included in the application only when building with the `copilot_inprocess` build tag. +The bundled runtime also provides the native library required by the [in-process transport](#in-process-transport-experimental). The `copilot_inprocess` build tag enables the transport implementation; both connection modes use the same embedded runtime artifacts. ## In-process transport (Experimental) diff --git a/go/cmd/bundler/main.go b/go/cmd/bundler/main.go index 7ee078eea0..5ec43aae56 100644 --- a/go/cmd/bundler/main.go +++ b/go/cmd/bundler/main.go @@ -16,6 +16,7 @@ import ( "compress/gzip" "crypto/sha256" "encoding/base64" + "encoding/hex" "encoding/json" "flag" "fmt" @@ -30,34 +31,42 @@ import ( "regexp" "runtime" "strings" + "time" "github.com/klauspost/compress/zstd" ) const ( // Keep these URLs centralized so reviewers can verify all outbound calls in one place. - sdkModule = "github.com/github/copilot-sdk/go" - packageJSONURLFmt = "https://raw.githubusercontent.com/github/copilot-sdk/%s/nodejs/package.json" - packageLockURLFmt = "https://raw.githubusercontent.com/github/copilot-sdk/%s/nodejs/package-lock.json" - tarballURLFmt = "https://registry.npmjs.org/@github/copilot-%s/-/copilot-%s-%s.tgz" - licenseTarballFmt = "https://registry.npmjs.org/@github/copilot/-/copilot-%s.tgz" - defaultPackageName = "main" + sdkModule = "github.com/github/copilot-sdk/go" + packageJSONURLFmt = "https://raw.githubusercontent.com/github/copilot-sdk/%s/nodejs/package.json" + packageLockURLFmt = "https://raw.githubusercontent.com/github/copilot-sdk/%s/nodejs/package-lock.json" + defaultReleaseDownloadURL = "https://github.com/github/copilot-cli/releases/download" + releaseDownloadURLEnv = "COPILOT_CLI_DOWNLOAD_BASE_URL" + defaultPackageName = "main" + maxChecksumManifestSize = 1 << 20 + maxReleaseDownloadAttempts = 3 + bundleMetadataSchema = 1 ) -// Platform info: npm package suffix, binary name +var releaseHTTPClient = &http.Client{Timeout: 60 * time.Second} +var releaseRetryDelay = time.Sleep +var cliVersionPattern = regexp.MustCompile(`^\d+\.\d+\.\d+(?:-[0-9A-Za-z._-]+)?$`) + +// Platform info: CLI release classifier and binary name. type platformInfo struct { - npmPlatform string - binaryName string + releasePlatform string + binaryName string } -// Map from GOOS/GOARCH to npm platform info +// Map from GOOS/GOARCH to CLI release platform info. var platforms = map[string]platformInfo{ - "linux/amd64": {npmPlatform: "linux-x64", binaryName: "copilot"}, - "linux/arm64": {npmPlatform: "linux-arm64", binaryName: "copilot"}, - "darwin/amd64": {npmPlatform: "darwin-x64", binaryName: "copilot"}, - "darwin/arm64": {npmPlatform: "darwin-arm64", binaryName: "copilot"}, - "windows/amd64": {npmPlatform: "win32-x64", binaryName: "copilot.exe"}, - "windows/arm64": {npmPlatform: "win32-arm64", binaryName: "copilot.exe"}, + "linux/amd64": {releasePlatform: "linux-x64", binaryName: "copilot"}, + "linux/arm64": {releasePlatform: "linux-arm64", binaryName: "copilot"}, + "darwin/amd64": {releasePlatform: "darwin-x64", binaryName: "copilot"}, + "darwin/arm64": {releasePlatform: "darwin-arm64", binaryName: "copilot"}, + "windows/amd64": {releasePlatform: "win32-x64", binaryName: "copilot.exe"}, + "windows/arm64": {releasePlatform: "win32-arm64", binaryName: "copilot.exe"}, } // main is the CLI entry point. @@ -70,7 +79,7 @@ func main() { // Resolve version first so the default output name can include it. version := resolveCLIVersion(*cliVersion) - // Resolve platform once to validate input and get the npm package mapping. + // Resolve platform once to validate input and get the release classifier. goos, goarch, info, err := resolvePlatform(*platform) if err != nil { fmt.Fprintf(os.Stderr, "Error: %v\n", err) @@ -101,7 +110,7 @@ func main() { fmt.Printf("Building bundle for %s (CLI version %s)\n", *platform, version) - bundle, err := buildBundle(info, version, outputPath, goos) + bundle, err := buildBundle(info, version, outputPath, goos, true) if err != nil { fmt.Fprintf(os.Stderr, "Error: %v\n", err) os.Exit(1) @@ -110,8 +119,8 @@ func main() { var muslBundle bundleArtifacts if goos == "linux" { muslInfo := platformInfo{ - npmPlatform: strings.Replace(info.npmPlatform, "linux-", "linuxmusl-", 1), - binaryName: info.binaryName, + releasePlatform: strings.Replace(info.releasePlatform, "linux-", "linuxmusl-", 1), + binaryName: info.binaryName, } muslOutputPath := filepath.Join(*output, defaultOutputFileName(version, "linuxmusl", goarch, info.binaryName)) muslBundle, err = buildBundle( @@ -119,17 +128,13 @@ func main() { version, muslOutputPath, goos, + false, ) if err != nil { fmt.Fprintf(os.Stderr, "Error: %v\n", err) os.Exit(1) } } - if err := downloadCLILicense(version, outputPath); err != nil { - fmt.Fprintf(os.Stderr, "Error: failed to download CLI license: %v\n", err) - os.Exit(1) - } - // Generate the Go file with embed directive if err := generateGoFile( goos, @@ -178,19 +183,31 @@ func resolvePlatform(platform string) (string, string, platformInfo, error) { // resolveCLIVersion determines the CLI version from the flag or repo metadata. func resolveCLIVersion(flagValue string) string { - if flagValue != "" { - return flagValue + version := flagValue + if version == "" { + detectedVersion, err := detectCLIVersion() + if err != nil { + fmt.Fprintf(os.Stderr, "Error detecting CLI version: %v\n", err) + fmt.Fprintln(os.Stderr, "Hint: specify --cli-version explicitly, or run from a Go module that depends on github.com/github/copilot-sdk/go") + os.Exit(1) + } + version = detectedVersion + fmt.Printf("Auto-detected CLI version: %s\n", version) } - version, err := detectCLIVersion() - if err != nil { - fmt.Fprintf(os.Stderr, "Error detecting CLI version: %v\n", err) - fmt.Fprintln(os.Stderr, "Hint: specify --cli-version explicitly, or run from a Go module that depends on github.com/github/copilot-sdk/go") + if err := validateCLIVersion(version); err != nil { + fmt.Fprintf(os.Stderr, "Error: %v\n", err) os.Exit(1) } - fmt.Printf("Auto-detected CLI version: %s\n", version) return version } +func validateCLIVersion(version string) error { + if !cliVersionPattern.MatchString(version) { + return fmt.Errorf("invalid CLI version %q", version) + } + return nil +} + // defaultOutputFileName builds the default bundle filename for a platform. func defaultOutputFileName(version, goos, goarch, binaryName string) string { base := strings.TrimSuffix(binaryName, filepath.Ext(binaryName)) @@ -386,6 +403,24 @@ func isHex(s string) bool { return true } +func isSHA256(s string) bool { + return len(s) == sha256.Size*2 && isHex(s) +} + +func findReleaseChecksum(contents, assetName string) (string, error) { + for line := range strings.SplitSeq(contents, "\n") { + fields := strings.Fields(line) + if len(fields) < 2 || strings.TrimPrefix(fields[1], "*") != assetName { + continue + } + if !isSHA256(fields[0]) { + return "", fmt.Errorf("invalid SHA-256 for %s", assetName) + } + return strings.ToLower(fields[0]), nil + } + return "", fmt.Errorf("SHA256SUMS.txt does not contain %s", assetName) +} + type bundleArtifacts struct { binaryPath string binaryHash []byte @@ -397,36 +432,39 @@ type bundleArtifacts struct { assetsHash []byte } +type bundleMetadata struct { + Schema int `json:"schema"` + CLIVersion string `json:"cliVersion"` + Platform string `json:"platform"` + ReleaseAsset string `json:"releaseAsset"` + ReleaseHash string `json:"releaseHash"` + BinaryHash string `json:"binaryHash"` + RuntimeHash string `json:"runtimeHash"` + WrapperHash string `json:"wrapperHash"` + AssetsHash string `json:"assetsHash"` + LicenseHash string `json:"licenseHash,omitempty"` +} + // buildBundle downloads the CLI and native runtime artifacts from one platform package. -func buildBundle(info platformInfo, cliVersion, outputPath, goos string) (bundleArtifacts, error) { +func buildBundle(info platformInfo, cliVersion, outputPath, goos string, includeLicense bool) (bundleArtifacts, error) { outputDir := filepath.Dir(outputPath) if outputDir == "" { outputDir = "." } - runtimeArtifactPath := filepath.Join(outputDir, runtimeLibArtifactName(cliVersion, info.npmPlatform, goos)) - wrapperArtifactPath := filepath.Join(outputDir, runtimeWrapperArtifactName(cliVersion, info.npmPlatform, info.binaryName)) - assetsArtifactPath := filepath.Join(outputDir, runtimeAssetsArtifactName(cliVersion, info.npmPlatform)) + runtimeArtifactPath := filepath.Join(outputDir, runtimeLibArtifactName(cliVersion, info.releasePlatform, goos)) + wrapperArtifactPath := filepath.Join(outputDir, runtimeWrapperArtifactName(cliVersion, info.releasePlatform, info.binaryName)) + assetsArtifactPath := filepath.Join(outputDir, runtimeAssetsArtifactName(cliVersion, info.releasePlatform)) + artifacts := bundleArtifacts{ + binaryPath: outputPath, + runtimeArtifactPath: runtimeArtifactPath, + wrapperArtifactPath: wrapperArtifactPath, + assetsArtifactPath: assetsArtifactPath, + } - if filesExist(outputPath, runtimeArtifactPath, wrapperArtifactPath, assetsArtifactPath) { + if cached, ok := loadCachedBundle(artifacts, cliVersion, info.releasePlatform, includeLicense); ok { // Idempotent output avoids re-downloading in CI or local rebuilds. - fmt.Printf("Output runtime bundle for %s already exists, skipping download\n", info.npmPlatform) - binaryHash, err := sha256FileFromCompressed(outputPath) - if err != nil { - return bundleArtifacts{}, fmt.Errorf("failed to hash existing output: %w", err) - } - runtimeHash, err := sha256FileFromCompressed(runtimeArtifactPath) - if err != nil { - return bundleArtifacts{}, fmt.Errorf("failed to hash existing runtime.node: %w", err) - } - wrapperHash, err := sha256FileFromCompressed(wrapperArtifactPath) - if err != nil { - return bundleArtifacts{}, fmt.Errorf("failed to hash existing runtime wrapper: %w", err) - } - assetsHash, err := sha256File(assetsArtifactPath) - if err != nil { - return bundleArtifacts{}, fmt.Errorf("failed to hash existing runtime assets: %w", err) - } - return bundleArtifacts{outputPath, binaryHash, runtimeArtifactPath, runtimeHash, wrapperArtifactPath, wrapperHash, assetsArtifactPath, assetsHash}, nil + fmt.Printf("Output runtime bundle for %s already exists, skipping download\n", info.releasePlatform) + return cached, nil } // Create temp directory for download @@ -436,7 +474,7 @@ func buildBundle(info platformInfo, cliVersion, outputPath, goos string) (bundle } defer os.RemoveAll(tempDir) - binaryPath, tarballPath, err := downloadCLIBinary(info.npmPlatform, info.binaryName, cliVersion, tempDir) + binaryPath, tarballPath, releaseHash, err := downloadCLIBinary(info.releasePlatform, info.binaryName, cliVersion, tempDir) if err != nil { return bundleArtifacts{}, fmt.Errorf("failed to download CLI binary: %w", err) } @@ -446,6 +484,11 @@ func buildBundle(info platformInfo, cliVersion, outputPath, goos string) (bundle return bundleArtifacts{}, fmt.Errorf("failed to create output directory: %w", err) } } + if includeLicense { + if err := extractCLILicense(tarballPath, outputPath); err != nil { + return bundleArtifacts{}, fmt.Errorf("failed to extract CLI license: %w", err) + } + } binaryHash, err := sha256File(binaryPath) if err != nil { @@ -459,10 +502,10 @@ func buildBundle(info platformInfo, cliVersion, outputPath, goos string) (bundle if err := extractFileFromTarball( tarballPath, tempDir, - "package/prebuilds/"+info.npmPlatform+"/runtime.node", + "package/prebuilds/"+info.releasePlatform+"/runtime.node", "runtime.node", ); err != nil { - return bundleArtifacts{}, fmt.Errorf("runtime package is missing prebuilds/%s/runtime.node: %w", info.npmPlatform, err) + return bundleArtifacts{}, fmt.Errorf("runtime package is missing prebuilds/%s/runtime.node: %w", info.releasePlatform, err) } runtimeHash, err := sha256File(rawLibPath) if err != nil { @@ -477,10 +520,10 @@ func buildBundle(info platformInfo, cliVersion, outputPath, goos string) (bundle if err := extractFileFromTarball( tarballPath, tempDir, - "package/prebuilds/"+info.npmPlatform+"/"+wrapperName, + "package/prebuilds/"+info.releasePlatform+"/"+wrapperName, wrapperName, ); err != nil { - return bundleArtifacts{}, fmt.Errorf("runtime package is missing prebuilds/%s/%s: %w", info.npmPlatform, wrapperName, err) + return bundleArtifacts{}, fmt.Errorf("runtime package is missing prebuilds/%s/%s: %w", info.releasePlatform, wrapperName, err) } wrapperHash, err := sha256File(rawWrapperPath) if err != nil { @@ -496,12 +539,110 @@ func buildBundle(info platformInfo, cliVersion, outputPath, goos string) (bundle if err != nil { return bundleArtifacts{}, fmt.Errorf("failed to hash runtime assets: %w", err) } + artifacts.binaryHash = binaryHash + artifacts.runtimeHash = runtimeHash + artifacts.wrapperHash = wrapperHash + artifacts.assetsHash = assetsHash + if err := writeBundleMetadata(artifacts, cliVersion, info.releasePlatform, releaseHash, includeLicense); err != nil { + return bundleArtifacts{}, fmt.Errorf("failed to write bundle metadata: %w", err) + } fmt.Printf("Successfully created %s\n", outputPath) fmt.Printf("Successfully created %s\n", runtimeArtifactPath) fmt.Printf("Successfully created %s\n", wrapperArtifactPath) fmt.Printf("Successfully created %s\n", assetsArtifactPath) - return bundleArtifacts{outputPath, binaryHash, runtimeArtifactPath, runtimeHash, wrapperArtifactPath, wrapperHash, assetsArtifactPath, assetsHash}, nil + return artifacts, nil +} + +func loadCachedBundle(artifacts bundleArtifacts, cliVersion, platform string, includeLicense bool) (bundleArtifacts, bool) { + requiredPaths := []string{ + artifacts.binaryPath, + artifacts.runtimeArtifactPath, + artifacts.wrapperArtifactPath, + artifacts.assetsArtifactPath, + bundleMetadataPath(artifacts.binaryPath), + } + if includeLicense { + requiredPaths = append(requiredPaths, licensePathForOutput(artifacts.binaryPath)) + } + if !filesExist(requiredPaths...) { + return bundleArtifacts{}, false + } + + contents, err := os.ReadFile(bundleMetadataPath(artifacts.binaryPath)) + if err != nil { + return bundleArtifacts{}, false + } + var metadata bundleMetadata + if err := json.Unmarshal(contents, &metadata); err != nil || + metadata.Schema != bundleMetadataSchema || + metadata.CLIVersion != cliVersion || + metadata.Platform != platform || + metadata.ReleaseAsset != releaseAssetName(cliVersion, platform) || + !isSHA256(metadata.ReleaseHash) || + (includeLicense && !isSHA256(metadata.LicenseHash)) { + return bundleArtifacts{}, false + } + + binaryHash, err := sha256FileFromCompressed(artifacts.binaryPath) + if err != nil || hex.EncodeToString(binaryHash) != metadata.BinaryHash { + return bundleArtifacts{}, false + } + runtimeHash, err := sha256FileFromCompressed(artifacts.runtimeArtifactPath) + if err != nil || hex.EncodeToString(runtimeHash) != metadata.RuntimeHash { + return bundleArtifacts{}, false + } + wrapperHash, err := sha256FileFromCompressed(artifacts.wrapperArtifactPath) + if err != nil || hex.EncodeToString(wrapperHash) != metadata.WrapperHash { + return bundleArtifacts{}, false + } + assetsHash, err := sha256File(artifacts.assetsArtifactPath) + if err != nil || hex.EncodeToString(assetsHash) != metadata.AssetsHash { + return bundleArtifacts{}, false + } + if includeLicense { + licenseHash, err := sha256File(licensePathForOutput(artifacts.binaryPath)) + if err != nil || hex.EncodeToString(licenseHash) != metadata.LicenseHash { + return bundleArtifacts{}, false + } + } + + artifacts.binaryHash = binaryHash + artifacts.runtimeHash = runtimeHash + artifacts.wrapperHash = wrapperHash + artifacts.assetsHash = assetsHash + return artifacts, true +} + +func writeBundleMetadata(artifacts bundleArtifacts, cliVersion, platform, releaseHash string, includeLicense bool) error { + metadata := bundleMetadata{ + Schema: bundleMetadataSchema, + CLIVersion: cliVersion, + Platform: platform, + ReleaseAsset: releaseAssetName(cliVersion, platform), + ReleaseHash: releaseHash, + BinaryHash: hex.EncodeToString(artifacts.binaryHash), + RuntimeHash: hex.EncodeToString(artifacts.runtimeHash), + WrapperHash: hex.EncodeToString(artifacts.wrapperHash), + AssetsHash: hex.EncodeToString(artifacts.assetsHash), + } + if includeLicense { + licenseHash, err := sha256File(licensePathForOutput(artifacts.binaryPath)) + if err != nil { + return err + } + metadata.LicenseHash = hex.EncodeToString(licenseHash) + } + contents, err := json.MarshalIndent(metadata, "", " ") + if err != nil { + return err + } + contents = append(contents, '\n') + return os.WriteFile(bundleMetadataPath(artifacts.binaryPath), contents, 0644) +} + +func bundleMetadataPath(outputPath string) string { + return outputPath + ".bundle.json" } func filesExist(paths ...string) bool { @@ -514,16 +655,16 @@ func filesExist(paths ...string) bool { } // runtimeLibArtifactName builds the compressed runtime-library artifact filename. -func runtimeLibArtifactName(version, npmPlatform, goos string) string { - return fmt.Sprintf("zcopilotruntime_%s_%s.%s.zst", version, npmPlatform, runtimeLibExt(goos)) +func runtimeLibArtifactName(version, releasePlatform, goos string) string { + return fmt.Sprintf("zcopilotruntime_%s_%s.%s.zst", version, releasePlatform, runtimeLibExt(goos)) } -func runtimeWrapperArtifactName(version, npmPlatform, binaryName string) string { - return fmt.Sprintf("zcopilotruntimewrapper_%s_%s_%s.zst", version, npmPlatform, runtimeWrapperName(binaryName)) +func runtimeWrapperArtifactName(version, releasePlatform, binaryName string) string { + return fmt.Sprintf("zcopilotruntimewrapper_%s_%s_%s.zst", version, releasePlatform, runtimeWrapperName(binaryName)) } -func runtimeAssetsArtifactName(version, npmPlatform string) string { - return fmt.Sprintf("zcopilotruntimeassets_%s_%s.tgz", version, npmPlatform) +func runtimeAssetsArtifactName(version, releasePlatform string) string { + return fmt.Sprintf("zcopilotruntimeassets_%s_%s.tgz", version, releasePlatform) } func runtimeWrapperName(binaryName string) string { @@ -540,12 +681,20 @@ var hostlessExcludedTopLevel = map[string]bool{ "sea-loader.js": true, "webview": true, } -func hostlessRuntimePath(name, npmPlatform, wrapperName string) (string, bool) { +func hostlessRuntimePath(name, releasePlatform, wrapperName string) (string, bool) { + if strings.Contains(name, `\`) { + return "", false + } relative, ok := strings.CutPrefix(name, "package/") - if !ok { + if !ok || relative == "" { return "", false } parts := strings.Split(relative, "/") + for _, part := range parts { + if part == "" || part == "." || part == ".." || strings.Contains(part, ":") { + return "", false + } + } topLevel := parts[0] fileName := parts[len(parts)-1] if hostlessExcludedTopLevel[topLevel] || @@ -561,7 +710,7 @@ func hostlessRuntimePath(name, npmPlatform, wrapperName string) (string, bool) { } } if topLevel == "prebuilds" { - if len(parts) < 3 || parts[1] != npmPlatform { + if len(parts) < 3 || parts[1] != releasePlatform { return "", false } return strings.Join(parts[2:], "/"), true @@ -602,7 +751,7 @@ func createRuntimeAssetsArchive(tarballPath, outputPath string, info platformInf } destination, include := hostlessRuntimePath( header.Name, - info.npmPlatform, + info.releasePlatform, runtimeWrapperName(info.binaryName), ) if !include { @@ -641,9 +790,8 @@ func runtimeLibExt(goos string) string { } } -// generateGoFile creates separate source files for normal and in-process builds. -// Both embed the CLI, while only the copilot_inprocess-tagged file embeds the -// native runtime library. +// generateGoFile creates one platform-specific source file containing the CLI +// and runtime artifacts used by both managed and in-process connections. func generateGoFile( goos, goarch, @@ -674,9 +822,8 @@ func generateGoFile( } outputDir := filepath.Dir(binaryPath) - defaultPath := filepath.Join(outputDir, fmt.Sprintf("zcopilot_%s_%s.go", goos, goarch)) - defaultContent := generatedGoFileContent( - "!copilot_inprocess", + sourcePath := filepath.Join(outputDir, fmt.Sprintf("zcopilot_%s_%s.go", goos, goarch)) + content := generatedGoFileContent( pkgName, binaryName, licenseName, @@ -697,44 +844,20 @@ func generateGoFile( muslAssetsArtifactPath, muslAssetsHash, ) - if err := os.WriteFile(defaultPath, []byte(defaultContent), 0644); err != nil { + if err := os.WriteFile(sourcePath, []byte(content), 0644); err != nil { return err } - inProcessPath := filepath.Join(outputDir, fmt.Sprintf("zcopilot_inprocess_%s_%s.go", goos, goarch)) - inProcessContent := generatedGoFileContent( - "copilot_inprocess", - pkgName, - binaryName, - licenseName, - cliVersion, - hashBase64, - runtimeArtifactPath, - runtimeHash, - wrapperArtifactPath, - wrapperHash, - assetsArtifactPath, - assetsHash, - muslBinaryPath, - muslBinaryHash, - muslRuntimeArtifactPath, - muslRuntimeHash, - muslWrapperArtifactPath, - muslWrapperHash, - muslAssetsArtifactPath, - muslAssetsHash, - ) - if err := os.WriteFile(inProcessPath, []byte(inProcessContent), 0644); err != nil { + legacyInProcessPath := filepath.Join(outputDir, fmt.Sprintf("zcopilot_inprocess_%s_%s.go", goos, goarch)) + if err := os.Remove(legacyInProcessPath); err != nil && !os.IsNotExist(err) { return err } - fmt.Printf("Generated %s\n", defaultPath) - fmt.Printf("Generated %s\n", inProcessPath) + fmt.Printf("Generated %s\n", sourcePath) return nil } func generatedGoFileContent( - buildConstraint, pkgName, binaryName, licenseName, @@ -757,7 +880,6 @@ func generatedGoFileContent( ) string { runtimeEmbed := "" runtimeConfig := "" - runtimeReader := "" if runtimeArtifactPath != "" && wrapperArtifactPath != "" && assetsArtifactPath != "" { runtimeArtifactName := filepath.Base(runtimeArtifactPath) runtimeHashBase64 := base64.StdEncoding.EncodeToString(runtimeHash) @@ -776,36 +898,18 @@ var localEmbeddedCopilotRuntimeExecutable []byte var localEmbeddedCopilotRuntimeAssets []byte `, runtimeArtifactName, wrapperArtifactName, assetsArtifactName) runtimeConfig = fmt.Sprintf(` - RuntimeLib: runtimeLibReader(), + RuntimeLib: zstdReader(localEmbeddedCopilotRuntimeLib), RuntimeLibHash: mustDecodeBase64(%q), - RuntimeNode: runtimeLibReader(), + RuntimeNode: zstdReader(localEmbeddedCopilotRuntimeLib), RuntimeNodeHash: mustDecodeBase64(%q), - RuntimeExecutable: runtimeExecutableReader(), + RuntimeExecutable: zstdReader(localEmbeddedCopilotRuntimeExecutable), RuntimeExecutableHash: mustDecodeBase64(%q), RuntimeAssets: bytes.NewReader(localEmbeddedCopilotRuntimeAssets), RuntimeAssetsHash: mustDecodeBase64(%q),`, runtimeHashBase64, runtimeHashBase64, wrapperHashBase64, assetsHashBase64) - runtimeReader = ` -func runtimeLibReader() io.Reader { - r, err := zstd.NewReader(bytes.NewReader(localEmbeddedCopilotRuntimeLib)) - if err != nil { - panic("failed to create zstd reader: " + err.Error()) - } - return r -} - -func runtimeExecutableReader() io.Reader { - r, err := zstd.NewReader(bytes.NewReader(localEmbeddedCopilotRuntimeExecutable)) - if err != nil { - panic("failed to create zstd reader: " + err.Error()) - } - return r -} -` } muslEmbed := "" muslConfig := "" - muslReaders := "" if muslBinaryPath != "" && muslRuntimeArtifactPath != "" && muslWrapperArtifactPath != "" && muslAssetsArtifactPath != "" { muslBinaryName := filepath.Base(muslBinaryPath) muslBinaryHashBase64 := base64.StdEncoding.EncodeToString(muslBinaryHash) @@ -829,46 +933,19 @@ var localEmbeddedCopilotRuntimeExecutableLinuxMusl []byte var localEmbeddedCopilotRuntimeAssetsLinuxMusl []byte `, muslBinaryName, muslRuntimeName, muslWrapperName, muslAssetsName) muslConfig = fmt.Sprintf(` - LinuxMuslCli: linuxMuslCLIReader(), + LinuxMuslCli: zstdReader(localEmbeddedCopilotCLILinuxMusl), LinuxMuslCliHash: mustDecodeBase64(%q), - LinuxMuslRuntimeLib: linuxMuslRuntimeLibReader(), + LinuxMuslRuntimeLib: zstdReader(localEmbeddedCopilotRuntimeLibLinuxMusl), LinuxMuslRuntimeLibHash: mustDecodeBase64(%q), - LinuxMuslRuntimeNode: linuxMuslRuntimeLibReader(), + LinuxMuslRuntimeNode: zstdReader(localEmbeddedCopilotRuntimeLibLinuxMusl), LinuxMuslRuntimeNodeHash: mustDecodeBase64(%q), - LinuxMuslRuntimeExecutable: linuxMuslRuntimeExecutableReader(), + LinuxMuslRuntimeExecutable: zstdReader(localEmbeddedCopilotRuntimeExecutableLinuxMusl), LinuxMuslRuntimeExecutableHash: mustDecodeBase64(%q), LinuxMuslRuntimeAssets: bytes.NewReader(localEmbeddedCopilotRuntimeAssetsLinuxMusl), LinuxMuslRuntimeAssetsHash: mustDecodeBase64(%q),`, muslBinaryHashBase64, muslRuntimeHashBase64, muslRuntimeHashBase64, muslWrapperHashBase64, muslAssetsHashBase64) - muslReaders = ` -func linuxMuslCLIReader() io.Reader { - r, err := zstd.NewReader(bytes.NewReader(localEmbeddedCopilotCLILinuxMusl)) - if err != nil { - panic("failed to create zstd reader: " + err.Error()) - } - return r -} - -func linuxMuslRuntimeLibReader() io.Reader { - r, err := zstd.NewReader(bytes.NewReader(localEmbeddedCopilotRuntimeLibLinuxMusl)) - if err != nil { - panic("failed to create zstd reader: " + err.Error()) - } - return r -} - -func linuxMuslRuntimeExecutableReader() io.Reader { - r, err := zstd.NewReader(bytes.NewReader(localEmbeddedCopilotRuntimeExecutableLinuxMusl)) - if err != nil { - panic("failed to create zstd reader: " + err.Error()) - } - return r -} -` } - return fmt.Sprintf(`//go:build %s - -// Code generated by copilot-sdk bundler; DO NOT EDIT. + return fmt.Sprintf(`// Code generated by copilot-sdk bundler; DO NOT EDIT. package %s @@ -892,22 +969,20 @@ var localEmbeddedCopilotCLILicense []byte func init() { embeddedcli.Setup(embeddedcli.Config{ - Cli: cliReader(), + Cli: zstdReader(localEmbeddedCopilotCLI), License: localEmbeddedCopilotCLILicense, Version: %q, CliHash: mustDecodeBase64(%q),%s%s }) } -func cliReader() io.Reader { - r, err := zstd.NewReader(bytes.NewReader(localEmbeddedCopilotCLI)) +func zstdReader(data []byte) io.Reader { + r, err := zstd.NewReader(bytes.NewReader(data)) if err != nil { panic("failed to create zstd reader: " + err.Error()) } return r } -%s -%s func mustDecodeBase64(s string) []byte { b, err := base64.StdEncoding.DecodeString(s) if err != nil { @@ -915,93 +990,152 @@ func mustDecodeBase64(s string) []byte { } return b } -`, buildConstraint, pkgName, binaryName, licenseName, runtimeEmbed, muslEmbed, cliVersion, hashBase64, runtimeConfig, muslConfig, runtimeReader, muslReaders) +`, pkgName, binaryName, licenseName, runtimeEmbed, muslEmbed, cliVersion, hashBase64, runtimeConfig, muslConfig) } -// downloadCLIBinary downloads the npm tarball and extracts the CLI binary. It -// returns the extracted binary path and the downloaded tarball path (retained so -// callers can extract additional files, such as the runtime library). -func downloadCLIBinary(npmPlatform, binaryName, cliVersion, destDir string) (string, string, error) { - tarballURL := fmt.Sprintf(tarballURLFmt, npmPlatform, npmPlatform, cliVersion) +// downloadCLIBinary downloads and verifies the CLI release archive, then +// extracts the CLI binary. It returns the extracted binary path and archive path +// so callers can extract the runtime artifacts from the same verified archive. +func downloadCLIBinary(releasePlatform, binaryName, cliVersion, destDir string) (string, string, string, error) { + assetName := releaseAssetName(cliVersion, releasePlatform) + releaseURL := fmt.Sprintf("%s/v%s", releaseDownloadBaseURL(), cliVersion) + expectedChecksum, err := downloadReleaseChecksum(releaseURL, assetName) + if err != nil { + return "", "", "", err + } + tarballURL := releaseURL + "/" + assetName fmt.Printf("Downloading from %s...\n", tarballURL) - resp, err := http.Get(tarballURL) + resp, err := getReleaseURL(tarballURL) if err != nil { - return "", "", fmt.Errorf("failed to download: %w", err) + return "", "", "", err } defer resp.Body.Close() - if resp.StatusCode != http.StatusOK { - return "", "", fmt.Errorf("failed to download: %s", resp.Status) - } - // Save tarball to temp file - tarballPath := filepath.Join(destDir, fmt.Sprintf("copilot-%s-%s.tgz", npmPlatform, cliVersion)) + tarballPath := filepath.Join(destDir, assetName) tarballFile, err := os.Create(tarballPath) if err != nil { - return "", "", fmt.Errorf("failed to create tarball file: %w", err) + return "", "", "", fmt.Errorf("failed to create tarball file: %w", err) } - if _, err := io.Copy(tarballFile, resp.Body); err != nil { + hasher := sha256.New() + if _, err := io.Copy(io.MultiWriter(tarballFile, hasher), resp.Body); err != nil { tarballFile.Close() - return "", "", fmt.Errorf("failed to save tarball: %w", err) + return "", "", "", fmt.Errorf("failed to save tarball: %w", err) } if err := tarballFile.Close(); err != nil { - return "", "", fmt.Errorf("failed to close tarball file: %w", err) + return "", "", "", fmt.Errorf("failed to close tarball file: %w", err) + } + actualChecksum := hex.EncodeToString(hasher.Sum(nil)) + if actualChecksum != expectedChecksum { + _ = os.Remove(tarballPath) + return "", "", "", fmt.Errorf("checksum mismatch for %s: expected %s, got %s", assetName, expectedChecksum, actualChecksum) } + fmt.Printf("Integrity verified for %s\n", assetName) // Extract only the CLI binary to avoid unpacking the full package tree. binaryPath := filepath.Join(destDir, binaryName) if err := extractFileFromTarball(tarballPath, destDir, "package/"+binaryName, binaryName); err != nil { - return "", "", fmt.Errorf("failed to extract binary: %w", err) + return "", "", "", fmt.Errorf("failed to extract binary: %w", err) } // Verify binary exists if _, err := os.Stat(binaryPath); err != nil { - return "", "", fmt.Errorf("binary not found after extraction: %w", err) + return "", "", "", fmt.Errorf("binary not found after extraction: %w", err) } // Make executable on Unix if !strings.HasSuffix(binaryName, ".exe") { if err := os.Chmod(binaryPath, 0755); err != nil { - return "", "", fmt.Errorf("failed to chmod binary: %w", err) + return "", "", "", fmt.Errorf("failed to chmod binary: %w", err) } } stat, err := os.Stat(binaryPath) if err != nil { - return "", "", fmt.Errorf("failed to stat binary: %w", err) + return "", "", "", fmt.Errorf("failed to stat binary: %w", err) } sizeMB := float64(stat.Size()) / 1024 / 1024 fmt.Printf("Downloaded %s (%.1f MB)\n", binaryName, sizeMB) - return binaryPath, tarballPath, nil + return binaryPath, tarballPath, actualChecksum, nil } -// downloadCLILicense downloads the @github/copilot package and writes its license next to outputPath. -func downloadCLILicense(cliVersion, outputPath string) error { - outputDir := filepath.Dir(outputPath) - if outputDir == "" { - outputDir = "." +func releaseAssetName(cliVersion, platform string) string { + return fmt.Sprintf("github-copilot-%s-%s.tgz", cliVersion, platform) +} + +func releaseDownloadBaseURL() string { + if value := strings.TrimRight(os.Getenv(releaseDownloadURLEnv), "/"); value != "" { + return value } - licensePath := licensePathForOutput(outputPath) - if _, err := os.Stat(licensePath); err == nil { - return nil + return defaultReleaseDownloadURL +} + +func getReleaseURL(url string) (*http.Response, error) { + var lastErr error + for attempt := range maxReleaseDownloadAttempts { + resp, err := releaseHTTPClient.Get(url) + if err == nil { + if resp.StatusCode == http.StatusOK { + return resp, nil + } + status := resp.Status + _ = resp.Body.Close() + lastErr = fmt.Errorf("server returned %s", status) + if !isRetriableHTTPStatus(resp.StatusCode) { + return nil, fmt.Errorf("failed to download %s: %w", url, lastErr) + } + } else { + lastErr = err + } + + if attempt+1 < maxReleaseDownloadAttempts { + releaseRetryDelay(time.Duration(1<= 500 +} + +func downloadReleaseChecksum(releaseURL, assetName string) (string, error) { + checksumsURL := releaseURL + "/SHA256SUMS.txt" + resp, err := getReleaseURL(checksumsURL) if err != nil { - return fmt.Errorf("failed to download license tarball: %w", err) + return "", err } defer resp.Body.Close() - if resp.StatusCode != http.StatusOK { - return fmt.Errorf("failed to download license tarball: %s", resp.Status) + contents, err := io.ReadAll(io.LimitReader(resp.Body, maxChecksumManifestSize+1)) + if err != nil { + return "", fmt.Errorf("failed to read checksums: %w", err) + } + if len(contents) > maxChecksumManifestSize { + return "", fmt.Errorf("SHA256SUMS.txt exceeds %d bytes", maxChecksumManifestSize) + } + return findReleaseChecksum(string(contents), assetName) +} + +// extractCLILicense writes the license from a verified release archive next to outputPath. +func extractCLILicense(tarballPath, outputPath string) error { + outputDir := filepath.Dir(outputPath) + if outputDir == "" { + outputDir = "." } + licensePath := licensePathForOutput(outputPath) - gzReader, err := gzip.NewReader(resp.Body) + sourceFile, err := os.Open(tarballPath) + if err != nil { + return fmt.Errorf("failed to open release archive: %w", err) + } + defer sourceFile.Close() + + gzReader, err := gzip.NewReader(sourceFile) if err != nil { return fmt.Errorf("failed to create gzip reader: %w", err) } @@ -1030,15 +1164,15 @@ func downloadCLILicense(cliVersion, outputPath string) error { } func licensePathForOutput(outputPath string) string { - if strings.HasSuffix(outputPath, ".zst") { - return strings.TrimSuffix(outputPath, ".zst") + ".license" + if before, ok := strings.CutSuffix(outputPath, ".zst"); ok { + return before + ".license" } return outputPath + ".license" } func licenseFileName(binaryName string) string { - if strings.HasSuffix(binaryName, ".zst") { - return strings.TrimSuffix(binaryName, ".zst") + ".license" + if before, ok := strings.CutSuffix(binaryName, ".zst"); ok { + return before + ".license" } return binaryName + ".license" } @@ -1107,21 +1241,6 @@ func extractFileFromTarball(tarballPath, destDir, targetPath, outputName string) return fmt.Errorf("file %q not found in tarball", targetPath) } -// extractOptionalFileFromTarball extracts a single file from a .tgz into destDir -// like extractFileFromTarball, but returns (false, nil) instead of an error when -// the file is absent. Used for the runtime library, which older CLI packages do -// not ship. -func extractOptionalFileFromTarball(tarballPath, destDir, targetPath, outputName string) (bool, error) { - err := extractFileFromTarball(tarballPath, destDir, targetPath, outputName) - if err == nil { - return true, nil - } - if strings.Contains(err.Error(), "not found in tarball") { - return false, nil - } - return false, err -} - // compressZstdFile compresses src into dst using zstd. func compressZstdFile(src, dst string) error { srcFile, err := os.Open(src) diff --git a/go/cmd/bundler/main_test.go b/go/cmd/bundler/main_test.go index 9c09559b66..66a9cac136 100644 --- a/go/cmd/bundler/main_test.go +++ b/go/cmd/bundler/main_test.go @@ -4,15 +4,272 @@ import ( "archive/tar" "bytes" "compress/gzip" + "crypto/sha256" + "fmt" "go/parser" "go/token" "io" + "net/http" + "net/http/httptest" "os" "path/filepath" "strings" + "sync/atomic" "testing" + "time" ) +func TestFindReleaseChecksum(t *testing.T) { + expected := strings.Repeat("a", 64) + checksums := strings.Join([]string{ + strings.Repeat("b", 64) + " other.tgz", + expected + " *github-copilot-1.2.3-linux-x64.tgz", + }, "\n") + + got, err := findReleaseChecksum(checksums, "github-copilot-1.2.3-linux-x64.tgz") + if err != nil { + t.Fatal(err) + } + if got != expected { + t.Fatalf("findReleaseChecksum() = %q, want %q", got, expected) + } + + if _, err := findReleaseChecksum(checksums, "missing.tgz"); err == nil { + t.Fatal("findReleaseChecksum() succeeded for a missing asset") + } +} + +func TestReleaseAssetName(t *testing.T) { + for _, test := range []struct { + platform string + want string + }{ + {platform: "linux-x64", want: "github-copilot-1.2.3-linux-x64.tgz"}, + {platform: "linuxmusl-arm64", want: "github-copilot-1.2.3-linuxmusl-arm64.tgz"}, + {platform: "win32-x64", want: "github-copilot-1.2.3-win32-x64.tgz"}, + } { + t.Run(test.platform, func(t *testing.T) { + if got := releaseAssetName("1.2.3", test.platform); got != test.want { + t.Fatalf("releaseAssetName() = %q, want %q", got, test.want) + } + }) + } +} + +func TestValidateCLIVersion(t *testing.T) { + for _, version := range []string{"1.2.3", "1.2.3-4", "1.2.3-preview.1"} { + if err := validateCLIVersion(version); err != nil { + t.Errorf("validateCLIVersion(%q) returned %v", version, err) + } + } + for _, version := range []string{"", "v1.2.3", "1.2", "../1.2.3", "1.2.3/asset"} { + if err := validateCLIVersion(version); err == nil { + t.Errorf("validateCLIVersion(%q) succeeded", version) + } + } +} + +func TestGetReleaseURLRetriesTransientFailures(t *testing.T) { + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + if attempts.Add(1) == 1 { + http.Error(w, "try again", http.StatusServiceUnavailable) + return + } + _, _ = w.Write([]byte("ok")) + })) + defer server.Close() + + oldDelay := releaseRetryDelay + releaseRetryDelay = func(time.Duration) {} + defer func() { releaseRetryDelay = oldDelay }() + + resp, err := getReleaseURL(server.URL) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + if got := attempts.Load(); got != 2 { + t.Fatalf("attempts = %d, want 2", got) + } +} + +func TestGetReleaseURLDoesNotRetryPermanentFailure(t *testing.T) { + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempts.Add(1) + http.NotFound(w, r) + })) + defer server.Close() + + if _, err := getReleaseURL(server.URL); err == nil { + t.Fatal("getReleaseURL() succeeded for HTTP 404") + } + if got := attempts.Load(); got != 1 { + t.Fatalf("attempts = %d, want 1", got) + } +} + +func TestDownloadCLIBinaryUsesVerifiedReleaseAsset(t *testing.T) { + dir := t.TempDir() + archivePath := filepath.Join(dir, "source.tgz") + writeTarGz(t, archivePath, map[string]string{"package/copilot": "binary"}) + archive, err := os.ReadFile(archivePath) + if err != nil { + t.Fatal(err) + } + assetName := "github-copilot-1.2.3-linux-x64.tgz" + checksum := sha256.Sum256(archive) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/v1.2.3/SHA256SUMS.txt": + fmt.Fprintf(w, "%x %s\n", checksum, assetName) + case "/v1.2.3/" + assetName: + _, _ = w.Write(archive) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + t.Setenv(releaseDownloadURLEnv, server.URL+"/") + + binaryPath, downloadedArchivePath, archiveChecksum, err := downloadCLIBinary("linux-x64", "copilot", "1.2.3", t.TempDir()) + if err != nil { + t.Fatal(err) + } + binary, err := os.ReadFile(binaryPath) + if err != nil { + t.Fatal(err) + } + if string(binary) != "binary" { + t.Fatalf("binary contents = %q, want %q", binary, "binary") + } + if filepath.Base(downloadedArchivePath) != assetName { + t.Fatalf("archive name = %q, want %q", filepath.Base(downloadedArchivePath), assetName) + } + if archiveChecksum != fmt.Sprintf("%x", checksum) { + t.Fatalf("archive checksum = %q, want %x", archiveChecksum, checksum) + } +} + +func TestDownloadCLIBinaryRejectsChecksumMismatch(t *testing.T) { + dir := t.TempDir() + archivePath := filepath.Join(dir, "source.tgz") + writeTarGz(t, archivePath, map[string]string{"package/copilot": "binary"}) + archive, err := os.ReadFile(archivePath) + if err != nil { + t.Fatal(err) + } + assetName := "github-copilot-1.2.3-linux-x64.tgz" + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/v1.2.3/SHA256SUMS.txt": + fmt.Fprintf(w, "%s %s\n", strings.Repeat("0", 64), assetName) + case "/v1.2.3/" + assetName: + _, _ = w.Write(archive) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + t.Setenv(releaseDownloadURLEnv, server.URL) + + destination := t.TempDir() + _, downloadedArchivePath, _, err := downloadCLIBinary("linux-x64", "copilot", "1.2.3", destination) + if err == nil || !strings.Contains(err.Error(), "checksum mismatch") { + t.Fatalf("downloadCLIBinary() error = %v, want checksum mismatch", err) + } + if downloadedArchivePath != "" { + t.Fatalf("downloaded archive path = %q, want empty", downloadedArchivePath) + } + if _, err := os.Stat(filepath.Join(destination, assetName)); !os.IsNotExist(err) { + t.Fatalf("unverified archive was not removed: %v", err) + } +} + +func TestExtractCLILicenseFromReleaseArchive(t *testing.T) { + dir := t.TempDir() + archivePath := filepath.Join(dir, "release.tgz") + writeTarGz(t, archivePath, map[string]string{"package/LICENSE.md": "license text"}) + outputPath := filepath.Join(dir, "zcopilot.zst") + + if err := extractCLILicense(archivePath, outputPath); err != nil { + t.Fatal(err) + } + license, err := os.ReadFile(licensePathForOutput(outputPath)) + if err != nil { + t.Fatal(err) + } + if string(license) != "license text" { + t.Fatalf("license contents = %q, want %q", license, "license text") + } +} + +func TestBuildBundleRefreshesCorruptCache(t *testing.T) { + dir := t.TempDir() + archivePath := filepath.Join(dir, "source.tgz") + writeTarGz(t, archivePath, map[string]string{ + "package/copilot": "binary", + "package/LICENSE.md": "license", + "package/prebuilds/linux-x64/runtime.node": "runtime", + "package/prebuilds/linux-x64/copilot-runtime": "wrapper", + "package/copilot-sdk/extension.js": "extension", + }) + archive, err := os.ReadFile(archivePath) + if err != nil { + t.Fatal(err) + } + assetName := releaseAssetName("1.2.3", "linux-x64") + checksum := sha256.Sum256(archive) + var requests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests.Add(1) + switch r.URL.Path { + case "/v1.2.3/SHA256SUMS.txt": + fmt.Fprintf(w, "%x %s\n", checksum, assetName) + case "/v1.2.3/" + assetName: + _, _ = w.Write(archive) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + t.Setenv(releaseDownloadURLEnv, server.URL) + + outputPath := filepath.Join(t.TempDir(), "zcopilot.zst") + info := platformInfo{releasePlatform: "linux-x64", binaryName: "copilot"} + bundle, err := buildBundle(info, "1.2.3", outputPath, "linux", true) + if err != nil { + t.Fatal(err) + } + if _, err := buildBundle(info, "1.2.3", outputPath, "linux", true); err != nil { + t.Fatal(err) + } + if got := requests.Load(); got != 2 { + t.Fatalf("release requests = %d, want 2 when reusing valid cache", got) + } + if err := os.WriteFile(bundle.assetsArtifactPath, []byte("corrupt"), 0644); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(licensePathForOutput(outputPath), []byte("corrupt"), 0644); err != nil { + t.Fatal(err) + } + + if _, err := buildBundle(info, "1.2.3", outputPath, "linux", true); err != nil { + t.Fatal(err) + } + if got := requests.Load(); got != 4 { + t.Fatalf("release requests = %d, want 4 after rebuilding corrupt cache", got) + } + license, err := os.ReadFile(licensePathForOutput(outputPath)) + if err != nil { + t.Fatal(err) + } + if string(license) != "license" { + t.Fatalf("license contents = %q, want %q", license, "license") + } +} + func TestCreateRuntimeAssetsArchiveRetainsUnknownAssetsAndFiltersCLIContent(t *testing.T) { dir := t.TempDir() source := filepath.Join(dir, "package.tgz") @@ -31,8 +288,8 @@ func TestCreateRuntimeAssetsArchiveRetainsUnknownAssetsAndFiltersCLIContent(t *t }) if err := createRuntimeAssetsArchive(source, output, platformInfo{ - npmPlatform: "linux-x64", - binaryName: "copilot", + releasePlatform: "linux-x64", + binaryName: "copilot", }); err != nil { t.Fatal(err) } @@ -54,6 +311,22 @@ func TestCreateRuntimeAssetsArchiveRetainsUnknownAssetsAndFiltersCLIContent(t *t } } +func TestHostlessRuntimePathRejectsUnsafePaths(t *testing.T) { + for _, name := range []string{ + "package/../escape", + "package/./asset", + "package//asset", + "package/C:/asset", + `package/assets\..\escape`, + } { + t.Run(name, func(t *testing.T) { + if destination, ok := hostlessRuntimePath(name, "linux-x64", "copilot-runtime"); ok { + t.Fatalf("hostlessRuntimePath() = %q, true; want rejected", destination) + } + }) + } +} + func writeTarGz(t *testing.T, path string, files map[string]string) { t.Helper() var buffer bytes.Buffer @@ -174,6 +447,7 @@ func TestGenerateGoFileEmbedsRuntimeWrapperPair(t *testing.T) { muslRuntimePath := filepath.Join(dir, "runtime-musl.node.zst") muslWrapperPath := filepath.Join(dir, "copilot-runtime-musl.zst") muslAssetsPath := filepath.Join(dir, "runtime-assets-musl.tgz") + legacyInProcessPath := filepath.Join(dir, "zcopilot_inprocess_linux_amd64.go") for _, path := range []string{ binaryPath, licensePathForOutput(binaryPath), @@ -184,6 +458,7 @@ func TestGenerateGoFileEmbedsRuntimeWrapperPair(t *testing.T) { muslRuntimePath, muslWrapperPath, muslAssetsPath, + legacyInProcessPath, } { if err := os.WriteFile(path, []byte("test"), 0644); err != nil { t.Fatal(err) @@ -220,8 +495,8 @@ func TestGenerateGoFileEmbedsRuntimeWrapperPair(t *testing.T) { if err != nil { t.Fatal(err) } - if !strings.Contains(string(defaultSource), "//go:build !copilot_inprocess") { - t.Fatal("default embed file does not exclude copilot_inprocess builds") + if strings.Contains(string(defaultSource), "//go:build") { + t.Fatal("platform embed file contains an unnecessary build constraint") } if !strings.Contains(string(defaultSource), "localEmbeddedCopilotRuntimeExecutable") { t.Fatal("default embed file does not include the runtime wrapper") @@ -238,27 +513,19 @@ func TestGenerateGoFileEmbedsRuntimeWrapperPair(t *testing.T) { if !strings.Contains(string(defaultSource), "localEmbeddedCopilotRuntimeLibLinuxMusl") { t.Fatal("default embed file does not include the Linux musl runtime") } + if !strings.Contains(string(defaultSource), "func zstdReader(data []byte) io.Reader") { + t.Fatal("generated embed file does not define a shared zstd reader") + } + for _, obsolete := range []string{"func cliReader()", "func runtimeLibReader()", "func linuxMuslCLIReader()"} { + if strings.Contains(string(defaultSource), obsolete) { + t.Fatalf("generated embed file contains obsolete reader %q", obsolete) + } + } if _, err := parser.ParseFile(token.NewFileSet(), "zcopilot_linux_amd64.go", defaultSource, parser.AllErrors); err != nil { t.Fatalf("default generated source is invalid: %v", err) } - inProcessSource, err := os.ReadFile(filepath.Join(dir, "zcopilot_inprocess_linux_amd64.go")) - if err != nil { - t.Fatal(err) - } - if !strings.Contains(string(inProcessSource), "//go:build copilot_inprocess") { - t.Fatal("in-process embed file does not require the copilot_inprocess tag") - } - if !strings.Contains(string(inProcessSource), "localEmbeddedCopilotRuntimeLib") { - t.Fatal("in-process embed file does not include the native runtime") - } - if !strings.Contains(string(inProcessSource), "localEmbeddedCopilotCLILinuxMusl") { - t.Fatal("in-process embed file does not include the Linux musl CLI") - } - if !strings.Contains(string(inProcessSource), "localEmbeddedCopilotRuntimeLibLinuxMusl") { - t.Fatal("in-process embed file does not include the Linux musl runtime") - } - if _, err := parser.ParseFile(token.NewFileSet(), "zcopilot_inprocess_linux_amd64.go", inProcessSource, parser.AllErrors); err != nil { - t.Fatalf("in-process generated source is invalid: %v", err) + if _, err := os.Stat(legacyInProcessPath); !os.IsNotExist(err) { + t.Fatalf("legacy in-process embed file was not removed: %v", err) } } diff --git a/go/internal/embeddedcli/embeddedcli.go b/go/internal/embeddedcli/embeddedcli.go index 2535cf5f20..fb5b9a4fab 100644 --- a/go/internal/embeddedcli/embeddedcli.go +++ b/go/internal/embeddedcli/embeddedcli.go @@ -26,7 +26,7 @@ import ( // when provided, is written next to the installed binary. // // RuntimeExecutable and RuntimeNode form the adjacent out-of-process runtime -// pair. RuntimeAssets is a filtered npm package archive containing auxiliary +// pair. RuntimeAssets is a filtered CLI release archive containing auxiliary // binaries and resources. RuntimeLib is the same cdylib bytes installed under // the natural platform name for the optional in-process transport. type Config struct { @@ -234,21 +234,11 @@ func isMusl() bool { } func installAt(installDir string) (string, error) { - version := sanitizeVersion(config.Version) - if version != "" { - installDir = filepath.Join(installDir, version) - } - if linuxMuslBundle { - installDir = filepath.Join(installDir, "linuxmusl") - } - if err := os.MkdirAll(installDir, 0755); err != nil { - return "", fmt.Errorf("creating install directory: %w", err) - } - - // Best effort to prevent concurrent installs. - if release, _ := flock.Acquire(filepath.Join(installDir, ".copilot-cli.lock")); release != nil { - defer release() + installDir, release, err := prepareInstallDir(installDir) + if err != nil { + return "", err } + defer release() binaryName := "copilot" if runtime.GOOS == "windows" { @@ -256,40 +246,8 @@ func installAt(installDir string) (string, error) { } finalPath := filepath.Join(installDir, binaryName) - if _, err := os.Stat(finalPath); err == nil { - existingHash, err := hashFile(finalPath) - if err != nil { - return "", fmt.Errorf("hashing existing binary: %w", err) - } - if !bytes.Equal(existingHash, config.CliHash) { - return "", fmt.Errorf("existing binary hash mismatch") - } - if config.RuntimeLib != nil { - libPath, err := installRuntimeLib(installDir) - if err != nil { - return "", err - } - runtimeLibPath = libPath - } - if err := installRuntimeAssets(installDir); err != nil { - return "", err - } - return finalPath, nil - } - - f, err := os.OpenFile(finalPath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0755) - if err != nil { - return "", fmt.Errorf("creating binary file: %w", err) - } - _, err = io.Copy(f, config.Cli) - if err1 := f.Close(); err1 != nil && err == nil { - err = err1 - } - if closer, ok := config.Cli.(io.Closer); ok { - closer.Close() - } - if err != nil { - return "", fmt.Errorf("writing binary file: %w", err) + if err := installVerifiedFile(finalPath, config.Cli, config.CliHash, 0755, "binary"); err != nil { + return "", err } if len(config.License) > 0 { licensePath := finalPath + ".license" @@ -305,16 +263,33 @@ func installAt(installDir string) (string, error) { if err != nil { return "", err } - if err := installRuntimeAssets(installDir); err != nil { - return "", err - } runtimeLibPath = libPath } + if err := installRuntimeAssets(installDir); err != nil { + return "", err + } return finalPath, nil } func installRuntimeAt(installDir string) (string, error) { + installDir, release, err := prepareInstallDir(installDir) + if err != nil { + return "", err + } + defer release() + + path, err := installRuntimePair(installDir) + if err != nil { + return "", err + } + if err := installRuntimeAssets(installDir); err != nil { + return "", err + } + return path, nil +} + +func prepareInstallDir(installDir string) (string, func() error, error) { version := sanitizeVersion(config.Version) if version != "" { installDir = filepath.Join(installDir, version) @@ -323,20 +298,15 @@ func installRuntimeAt(installDir string) (string, error) { installDir = filepath.Join(installDir, "linuxmusl") } if err := os.MkdirAll(installDir, 0755); err != nil { - return "", fmt.Errorf("creating install directory: %w", err) + return "", nil, fmt.Errorf("creating install directory: %w", err) } - if release, _ := flock.Acquire(filepath.Join(installDir, ".copilot-cli.lock")); release != nil { - defer release() - } - path, err := installRuntimePair(installDir) - if err != nil { - return "", err - } - if err := installRuntimeAssets(installDir); err != nil { - return "", err + // Best effort to prevent concurrent installs across processes. + release, _ := flock.Acquire(filepath.Join(installDir, ".copilot-cli.lock")) + if release == nil { + release = func() error { return nil } } - return path, nil + return installDir, release, nil } func validateOptionalHash(reader io.Reader, hash []byte, name string) { @@ -449,7 +419,7 @@ func installVerifiedFile(path string, reader io.Reader, expectedHash []byte, mod return nil } - tmp, err := os.CreateTemp(filepath.Dir(path), ".copilot-runtime-pair-*.tmp") + tmp, err := os.CreateTemp(filepath.Dir(path), ".copilot-artifact-*.tmp") if err != nil { return fmt.Errorf("creating temporary %s: %w", label, err) } @@ -495,43 +465,8 @@ func installRuntimeLib(installDir string) (string, error) { return "", fmt.Errorf("RuntimeLibHash must be a SHA-256 hash (%d bytes), got %d bytes", sha256.Size, len(config.RuntimeLibHash)) } libPath := filepath.Join(installDir, naturalRuntimeLibName()) - - if _, err := os.Stat(libPath); err == nil { - existingHash, err := hashFile(libPath) - if err != nil { - return "", fmt.Errorf("hashing existing runtime library: %w", err) - } - if !bytes.Equal(existingHash, config.RuntimeLibHash) { - return "", fmt.Errorf("existing runtime library hash mismatch") - } - return libPath, nil - } - - // Write to a temp file in the same directory, verify, then atomically rename. - tmp, err := os.CreateTemp(installDir, ".copilot-runtime-*.tmp") - if err != nil { - return "", fmt.Errorf("creating temp runtime library: %w", err) - } - tmpPath := tmp.Name() - h := sha256.New() - _, err = io.Copy(io.MultiWriter(tmp, h), config.RuntimeLib) - if err1 := tmp.Close(); err1 != nil && err == nil { - err = err1 - } - if closer, ok := config.RuntimeLib.(io.Closer); ok { - closer.Close() - } - if err != nil { - os.Remove(tmpPath) - return "", fmt.Errorf("writing runtime library: %w", err) - } - if !bytes.Equal(h.Sum(nil), config.RuntimeLibHash) { - os.Remove(tmpPath) - return "", fmt.Errorf("runtime library hash mismatch") - } - if err := os.Rename(tmpPath, libPath); err != nil { - os.Remove(tmpPath) - return "", fmt.Errorf("installing runtime library: %w", err) + if err := installVerifiedFile(libPath, config.RuntimeLib, config.RuntimeLibHash, 0644, "runtime library"); err != nil { + return "", err } return libPath, nil } From c83c913ddda6bdb9346dba7824d589bb88bc6941 Mon Sep 17 00:00:00 2001 From: Quim Muntal Date: Fri, 4 Sep 2026 12:00:48 +0200 Subject: [PATCH 2/3] Close response body if it exists on error Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- go/cmd/bundler/main.go | 3 +++ 1 file changed, 3 insertions(+) diff --git a/go/cmd/bundler/main.go b/go/cmd/bundler/main.go index 5ec43aae56..9845931d62 100644 --- a/go/cmd/bundler/main.go +++ b/go/cmd/bundler/main.go @@ -1089,6 +1089,9 @@ func getReleaseURL(url string) (*http.Response, error) { return nil, fmt.Errorf("failed to download %s: %w", url, lastErr) } } else { + if resp != nil { + _ = resp.Body.Close() + } lastErr = err } From f3498fc95a4994ed2ad3deb21573945381460dc5 Mon Sep 17 00:00:00 2001 From: qmuntal Date: Fri, 4 Sep 2026 12:04:51 +0200 Subject: [PATCH 3/3] Remove unenforced release hash metadata --- go/cmd/bundler/main.go | 33 +++++++++++++++------------------ go/cmd/bundler/main_test.go | 14 +++++++++----- 2 files changed, 24 insertions(+), 23 deletions(-) diff --git a/go/cmd/bundler/main.go b/go/cmd/bundler/main.go index 9845931d62..70fea0c78b 100644 --- a/go/cmd/bundler/main.go +++ b/go/cmd/bundler/main.go @@ -437,7 +437,6 @@ type bundleMetadata struct { CLIVersion string `json:"cliVersion"` Platform string `json:"platform"` ReleaseAsset string `json:"releaseAsset"` - ReleaseHash string `json:"releaseHash"` BinaryHash string `json:"binaryHash"` RuntimeHash string `json:"runtimeHash"` WrapperHash string `json:"wrapperHash"` @@ -474,7 +473,7 @@ func buildBundle(info platformInfo, cliVersion, outputPath, goos string, include } defer os.RemoveAll(tempDir) - binaryPath, tarballPath, releaseHash, err := downloadCLIBinary(info.releasePlatform, info.binaryName, cliVersion, tempDir) + binaryPath, tarballPath, err := downloadCLIBinary(info.releasePlatform, info.binaryName, cliVersion, tempDir) if err != nil { return bundleArtifacts{}, fmt.Errorf("failed to download CLI binary: %w", err) } @@ -543,7 +542,7 @@ func buildBundle(info platformInfo, cliVersion, outputPath, goos string, include artifacts.runtimeHash = runtimeHash artifacts.wrapperHash = wrapperHash artifacts.assetsHash = assetsHash - if err := writeBundleMetadata(artifacts, cliVersion, info.releasePlatform, releaseHash, includeLicense); err != nil { + if err := writeBundleMetadata(artifacts, cliVersion, info.releasePlatform, includeLicense); err != nil { return bundleArtifacts{}, fmt.Errorf("failed to write bundle metadata: %w", err) } @@ -579,7 +578,6 @@ func loadCachedBundle(artifacts bundleArtifacts, cliVersion, platform string, in metadata.CLIVersion != cliVersion || metadata.Platform != platform || metadata.ReleaseAsset != releaseAssetName(cliVersion, platform) || - !isSHA256(metadata.ReleaseHash) || (includeLicense && !isSHA256(metadata.LicenseHash)) { return bundleArtifacts{}, false } @@ -614,13 +612,12 @@ func loadCachedBundle(artifacts bundleArtifacts, cliVersion, platform string, in return artifacts, true } -func writeBundleMetadata(artifacts bundleArtifacts, cliVersion, platform, releaseHash string, includeLicense bool) error { +func writeBundleMetadata(artifacts bundleArtifacts, cliVersion, platform string, includeLicense bool) error { metadata := bundleMetadata{ Schema: bundleMetadataSchema, CLIVersion: cliVersion, Platform: platform, ReleaseAsset: releaseAssetName(cliVersion, platform), - ReleaseHash: releaseHash, BinaryHash: hex.EncodeToString(artifacts.binaryHash), RuntimeHash: hex.EncodeToString(artifacts.runtimeHash), WrapperHash: hex.EncodeToString(artifacts.wrapperHash), @@ -996,12 +993,12 @@ func mustDecodeBase64(s string) []byte { // downloadCLIBinary downloads and verifies the CLI release archive, then // extracts the CLI binary. It returns the extracted binary path and archive path // so callers can extract the runtime artifacts from the same verified archive. -func downloadCLIBinary(releasePlatform, binaryName, cliVersion, destDir string) (string, string, string, error) { +func downloadCLIBinary(releasePlatform, binaryName, cliVersion, destDir string) (string, string, error) { assetName := releaseAssetName(cliVersion, releasePlatform) releaseURL := fmt.Sprintf("%s/v%s", releaseDownloadBaseURL(), cliVersion) expectedChecksum, err := downloadReleaseChecksum(releaseURL, assetName) if err != nil { - return "", "", "", err + return "", "", err } tarballURL := releaseURL + "/" + assetName @@ -1009,7 +1006,7 @@ func downloadCLIBinary(releasePlatform, binaryName, cliVersion, destDir string) resp, err := getReleaseURL(tarballURL) if err != nil { - return "", "", "", err + return "", "", err } defer resp.Body.Close() @@ -1017,50 +1014,50 @@ func downloadCLIBinary(releasePlatform, binaryName, cliVersion, destDir string) tarballPath := filepath.Join(destDir, assetName) tarballFile, err := os.Create(tarballPath) if err != nil { - return "", "", "", fmt.Errorf("failed to create tarball file: %w", err) + return "", "", fmt.Errorf("failed to create tarball file: %w", err) } hasher := sha256.New() if _, err := io.Copy(io.MultiWriter(tarballFile, hasher), resp.Body); err != nil { tarballFile.Close() - return "", "", "", fmt.Errorf("failed to save tarball: %w", err) + return "", "", fmt.Errorf("failed to save tarball: %w", err) } if err := tarballFile.Close(); err != nil { - return "", "", "", fmt.Errorf("failed to close tarball file: %w", err) + return "", "", fmt.Errorf("failed to close tarball file: %w", err) } actualChecksum := hex.EncodeToString(hasher.Sum(nil)) if actualChecksum != expectedChecksum { _ = os.Remove(tarballPath) - return "", "", "", fmt.Errorf("checksum mismatch for %s: expected %s, got %s", assetName, expectedChecksum, actualChecksum) + return "", "", fmt.Errorf("checksum mismatch for %s: expected %s, got %s", assetName, expectedChecksum, actualChecksum) } fmt.Printf("Integrity verified for %s\n", assetName) // Extract only the CLI binary to avoid unpacking the full package tree. binaryPath := filepath.Join(destDir, binaryName) if err := extractFileFromTarball(tarballPath, destDir, "package/"+binaryName, binaryName); err != nil { - return "", "", "", fmt.Errorf("failed to extract binary: %w", err) + return "", "", fmt.Errorf("failed to extract binary: %w", err) } // Verify binary exists if _, err := os.Stat(binaryPath); err != nil { - return "", "", "", fmt.Errorf("binary not found after extraction: %w", err) + return "", "", fmt.Errorf("binary not found after extraction: %w", err) } // Make executable on Unix if !strings.HasSuffix(binaryName, ".exe") { if err := os.Chmod(binaryPath, 0755); err != nil { - return "", "", "", fmt.Errorf("failed to chmod binary: %w", err) + return "", "", fmt.Errorf("failed to chmod binary: %w", err) } } stat, err := os.Stat(binaryPath) if err != nil { - return "", "", "", fmt.Errorf("failed to stat binary: %w", err) + return "", "", fmt.Errorf("failed to stat binary: %w", err) } sizeMB := float64(stat.Size()) / 1024 / 1024 fmt.Printf("Downloaded %s (%.1f MB)\n", binaryName, sizeMB) - return binaryPath, tarballPath, actualChecksum, nil + return binaryPath, tarballPath, nil } func releaseAssetName(cliVersion, platform string) string { diff --git a/go/cmd/bundler/main_test.go b/go/cmd/bundler/main_test.go index 66a9cac136..4455f47b75 100644 --- a/go/cmd/bundler/main_test.go +++ b/go/cmd/bundler/main_test.go @@ -133,7 +133,7 @@ func TestDownloadCLIBinaryUsesVerifiedReleaseAsset(t *testing.T) { defer server.Close() t.Setenv(releaseDownloadURLEnv, server.URL+"/") - binaryPath, downloadedArchivePath, archiveChecksum, err := downloadCLIBinary("linux-x64", "copilot", "1.2.3", t.TempDir()) + binaryPath, downloadedArchivePath, err := downloadCLIBinary("linux-x64", "copilot", "1.2.3", t.TempDir()) if err != nil { t.Fatal(err) } @@ -147,9 +147,6 @@ func TestDownloadCLIBinaryUsesVerifiedReleaseAsset(t *testing.T) { if filepath.Base(downloadedArchivePath) != assetName { t.Fatalf("archive name = %q, want %q", filepath.Base(downloadedArchivePath), assetName) } - if archiveChecksum != fmt.Sprintf("%x", checksum) { - t.Fatalf("archive checksum = %q, want %x", archiveChecksum, checksum) - } } func TestDownloadCLIBinaryRejectsChecksumMismatch(t *testing.T) { @@ -175,7 +172,7 @@ func TestDownloadCLIBinaryRejectsChecksumMismatch(t *testing.T) { t.Setenv(releaseDownloadURLEnv, server.URL) destination := t.TempDir() - _, downloadedArchivePath, _, err := downloadCLIBinary("linux-x64", "copilot", "1.2.3", destination) + _, downloadedArchivePath, err := downloadCLIBinary("linux-x64", "copilot", "1.2.3", destination) if err == nil || !strings.Contains(err.Error(), "checksum mismatch") { t.Fatalf("downloadCLIBinary() error = %v, want checksum mismatch", err) } @@ -242,6 +239,13 @@ func TestBuildBundleRefreshesCorruptCache(t *testing.T) { if err != nil { t.Fatal(err) } + metadata, err := os.ReadFile(bundleMetadataPath(outputPath)) + if err != nil { + t.Fatal(err) + } + if strings.Contains(string(metadata), "releaseHash") { + t.Fatal("bundle metadata contains unenforced releaseHash") + } if _, err := buildBundle(info, "1.2.3", outputPath, "linux", true); err != nil { t.Fatal(err) }