From bcca909367e400d17400e6b12aabd60cbb65a31c Mon Sep 17 00:00:00 2001 From: Mike Date: Fri, 18 Sep 2026 16:25:48 -0700 Subject: [PATCH 1/3] refactor(api): add rpc-go release discovery and archive handling Groundwork for Download RPC; nothing calls it yet. - List v3+ rpc-go releases from GitHub (newest five) or from a local // cache, sorted newest first - Match rpc-go's published builds, rpc_linux_.tar.gz and rpc_windows_.exe - Take the binary from the bare .exe or the tarball's single entry, capped at 200 MiB - Assemble the download zip holding the binaries and config.yaml - Add the package request/release DTOs and the asset download URL --- go.mod | 1 + go.sum | 8 +- internal/entity/dto/v1/package.go | 36 ++++ internal/entity/dto/v1/package_test.go | 90 +++++++++ internal/entity/github/release.go | 13 +- internal/usecase/packaging/archive.go | 152 ++++++++++++++ internal/usecase/packaging/archive_test.go | 148 ++++++++++++++ internal/usecase/packaging/github.go | 162 +++++++++++++++ internal/usecase/packaging/github_test.go | 218 +++++++++++++++++++++ 9 files changed, 818 insertions(+), 10 deletions(-) create mode 100644 internal/entity/dto/v1/package.go create mode 100644 internal/entity/dto/v1/package_test.go create mode 100644 internal/usecase/packaging/archive.go create mode 100644 internal/usecase/packaging/archive_test.go create mode 100644 internal/usecase/packaging/github.go create mode 100644 internal/usecase/packaging/github_test.go diff --git a/go.mod b/go.mod index 1e76269f6..18214a06f 100644 --- a/go.mod +++ b/go.mod @@ -29,6 +29,7 @@ require ( github.com/zsais/go-gin-prometheus v1.0.3 go.mongodb.org/mongo-driver/v2 v2.9.0 go.uber.org/mock v0.6.0 + golang.org/x/mod v0.40.0 golang.org/x/sys v0.48.0 gopkg.in/yaml.v2 v2.4.0 modernc.org/sqlite v1.58.0 diff --git a/go.sum b/go.sum index c18fb2992..1005117a4 100644 --- a/go.sum +++ b/go.sum @@ -298,8 +298,8 @@ golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5y golang.org/x/crypto v0.56.0 h1:GUh5Ii4J5jtcseSMiRqr1jXCNHoxjeV9Fmekc2oLy6Y= golang.org/x/crypto v0.56.0/go.mod h1:OMW5y6CY9l38uPLmxU6l6pwcXp1obtLo3e6gT7gQR2I= golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4= -golang.org/x/mod v0.38.0 h1:MECBjubtXD7yj4HrhIUcywNaGeNVUdfVnxmPajOk4yk= -golang.org/x/mod v0.38.0/go.mod h1:V6Xz0pq8TQ3dGqVQ1FVHuelZpAL0uNhSkk9ogYP3c40= +golang.org/x/mod v0.40.0 h1:hUv+3cXcdRHz08UmSiOob7sadHig73uo5bkXxQ/tvUs= +golang.org/x/mod v0.40.0/go.mod h1:0/weTWkPWGBikyTWAX3dkjVztMmBA5hM0DH6BElSupE= golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= @@ -331,8 +331,8 @@ golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc= -golang.org/x/tools v0.48.0 h1:3+hClM1aLL5mjMKm5ovokw9epgRXPuu2tILgismM6RE= -golang.org/x/tools v0.48.0/go.mod h1:08xX0orndb/F7jJxGDicx061tyd5pcMto75YMAXr6lk= +golang.org/x/tools v0.49.0 h1:3NI7VXzL9+1WZD52Dx2ttoPwD5DWrFGpl9mFZDlmisI= +golang.org/x/tools v0.49.0/go.mod h1:SJNXV9DBKT0UbdttsQjbfJlAE/q+y36++zo3uL3N0Oo= golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= diff --git a/internal/entity/dto/v1/package.go b/internal/entity/dto/v1/package.go new file mode 100644 index 000000000..914b1c29d --- /dev/null +++ b/internal/entity/dto/v1/package.go @@ -0,0 +1,36 @@ +package dto + +// PackageAuth selects how rpc-go authenticates to the server. Mode "none" +// embeds no credentials; they are supplied on the device instead. +type PackageAuth struct { + Mode string `json:"mode" binding:"required,oneof=token userpass none"` + Username string `json:"username" binding:"required_if=Mode userpass"` + Password string `json:"password" binding:"required_if=Mode userpass"` +} + +// PackageRequest is the body posted to POST /api/package. +type PackageRequest struct { + Command string `json:"command" binding:"required,oneof=activate deactivate"` + Version string `json:"version" binding:"required"` + OS string `json:"os" binding:"required"` // "windows", "linux", or "both" + Arch string `json:"arch" binding:"required"` + Auth PackageAuth `json:"auth" binding:"required"` + Profile string `json:"profile" binding:"required_if=Command activate"` + Domain string `json:"domain"` + TokenTTL string `json:"tokenTtl" binding:"omitempty,oneof=15m 1h 8h 24h"` + // ServerURL is the base URL rpc-go is pointed at, e.g. + // "https://console.example.com:8181". Empty falls back to the listen address. + ServerURL string `json:"serverUrl" binding:"omitempty,url"` +} + +// RPCAsset is one downloadable build for a release. +type RPCAsset struct { + OS string `json:"os"` + Arch string `json:"arch"` +} + +// RPCRelease is a single rpc-go release returned to the UI. +type RPCRelease struct { + Version string `json:"version"` + Assets []RPCAsset `json:"assets"` +} diff --git a/internal/entity/dto/v1/package_test.go b/internal/entity/dto/v1/package_test.go new file mode 100644 index 000000000..0513b33a9 --- /dev/null +++ b/internal/entity/dto/v1/package_test.go @@ -0,0 +1,90 @@ +package dto + +import ( + "testing" + + "github.com/go-playground/validator/v10" + "github.com/stretchr/testify/require" +) + +// A package built for "activate" without a profile would point rpc-go at a +// profile-export URL with an empty name, which the server answers with 404. +func TestPackageRequestProfileRequiredForActivate(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + command string + profile string + wantErr bool + }{ + {"activate without profile is invalid", "activate", "", true}, + {"activate with profile is valid", "activate", "profile1", false}, + {"deactivate without profile is valid", "deactivate", "", false}, + } + + for _, tt := range tests { + tt := tt + + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + req := PackageRequest{ + Command: tt.command, + Version: "v3.0.1", + OS: "linux", + Arch: "x86_64", + Auth: PackageAuth{Mode: "userpass", Username: "u", Password: "p"}, + Profile: tt.profile, + } + + // Gin binds with the "binding" tag, not validator's default "validate". + validate := validator.New() + validate.SetTagName("binding") + + err := validate.Struct(req) + if tt.wantErr { + require.Error(t, err) + } else { + require.NoError(t, err) + } + }) + } +} + +func TestPackageAuthModes(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + auth PackageAuth + wantErr bool + }{ + {"none", PackageAuth{Mode: "none"}, false}, + {"token", PackageAuth{Mode: "token"}, false}, + {"userpass", PackageAuth{Mode: "userpass", Username: "u", Password: "p"}, false}, + {"userpass without password", PackageAuth{Mode: "userpass", Username: "u"}, true}, + {"userpass without username", PackageAuth{Mode: "userpass", Password: "p"}, true}, + {"userpass without credentials", PackageAuth{Mode: "userpass"}, true}, + {"empty", PackageAuth{}, true}, + {"basic", PackageAuth{Mode: "basic"}, true}, + } + + for _, tt := range tests { + tt := tt + + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + validate := validator.New() + validate.SetTagName("binding") + + err := validate.Struct(tt.auth) + if tt.wantErr { + require.Error(t, err) + } else { + require.NoError(t, err) + } + }) + } +} diff --git a/internal/entity/github/release.go b/internal/entity/github/release.go index e28d7e1bd..8d07f204b 100644 --- a/internal/entity/github/release.go +++ b/internal/entity/github/release.go @@ -24,10 +24,11 @@ type Author struct { } type Asset struct { - URL string `json:"url"` - ID int `json:"id"` - Name string `json:"name"` - Label string `json:"label"` - State string `json:"state"` - ContentType string `json:"content_type"` + URL string `json:"url"` + ID int `json:"id"` + Name string `json:"name"` + Label string `json:"label"` + State string `json:"state"` + ContentType string `json:"content_type"` + BrowserDownloadURL string `json:"browser_download_url"` } diff --git a/internal/usecase/packaging/archive.go b/internal/usecase/packaging/archive.go new file mode 100644 index 000000000..4d893f12b --- /dev/null +++ b/internal/usecase/packaging/archive.go @@ -0,0 +1,152 @@ +package packaging + +import ( + "archive/tar" + "archive/zip" + "bytes" + "compress/gzip" + "context" + "errors" + "fmt" + "io" + "net/http" + "strings" + "time" +) + +const ( + binaryName = "rpc" + binaryNameWin = "rpc.exe" + configFileName = "config.yaml" + binaryFileMode = 0o755 + maxArchiveBytes = 200 << 20 // 200 MiB cap to guard against decompression bombs + httpTimeout = 30 * time.Second +) + +// httpClient is shared by the releases list and asset downloads. +var httpClient = &http.Client{Timeout: httpTimeout} //nolint:gochecknoglobals // shared so connections are reused + +var ( + // ErrBinaryNotFound indicates the archive did not hold the expected rpc binary. + ErrBinaryNotFound = errors.New("rpc binary not found in archive") + // ErrEntryTooLarge indicates an archive entry exceeded the size cap. + ErrEntryTooLarge = errors.New("archive entry exceeds size limit") +) + +// readLimited reads up to maxArchiveBytes from r, guarding against decompression bombs. +func readLimited(r io.Reader) ([]byte, error) { + var buf bytes.Buffer + + n, err := io.CopyN(&buf, r, maxArchiveBytes+1) + if err != nil && !errors.Is(err, io.EOF) { + return nil, fmt.Errorf("read entry: %w", err) + } + + if n > maxArchiveBytes { + return nil, ErrEntryTooLarge + } + + return buf.Bytes(), nil +} + +// extractBinary returns the packaged name and bytes of an rpc-go asset: a bare +// .exe, or a .tar.gz holding one file named after the asset. +func extractBinary(data []byte, assetName string) (name string, content []byte, err error) { + if strings.HasSuffix(assetName, ".exe") { + return binaryNameWin, data, nil + } + + gz, err := gzip.NewReader(bytes.NewReader(data)) + if err != nil { + return "", nil, fmt.Errorf("gzip: %w", err) + } + defer gz.Close() + + tr := tar.NewReader(gz) + + hdr, err := tr.Next() + if err != nil { + return "", nil, fmt.Errorf("tar: %w", err) + } + + if hdr.Typeflag != tar.TypeReg || hdr.Name != strings.TrimSuffix(assetName, ".tar.gz") { + return "", nil, fmt.Errorf("%w: %s", ErrBinaryNotFound, hdr.Name) + } + + content, err = readLimited(tr) + if err != nil { + return "", nil, err + } + + return binaryName, content, nil +} + +// zipEntry is one binary placed in the downloadable zip. +type zipEntry struct { + name string + content []byte +} + +// buildZip assembles the downloadable zip containing the binaries and config.yaml. +func buildZip(binaries []zipEntry, configYAML []byte) ([]byte, error) { + var buf bytes.Buffer + + size := len(configYAML) + for _, b := range binaries { + size += len(b.content) + } + + buf.Grow(size) + + zw := zip.NewWriter(&buf) + + for _, b := range binaries { + binHeader := &zip.FileHeader{Name: b.name, Method: zip.Deflate} + binHeader.SetMode(binaryFileMode) + + bw, err := zw.CreateHeader(binHeader) + if err != nil { + return nil, fmt.Errorf("zip create binary: %w", err) + } + + if _, err := bw.Write(b.content); err != nil { + return nil, fmt.Errorf("zip write binary: %w", err) + } + } + + cw, err := zw.Create(configFileName) + if err != nil { + return nil, fmt.Errorf("zip create config: %w", err) + } + + if _, err := cw.Write(configYAML); err != nil { + return nil, fmt.Errorf("zip write config: %w", err) + } + + if err := zw.Close(); err != nil { + return nil, fmt.Errorf("zip close: %w", err) + } + + return buf.Bytes(), nil +} + +// httpGet GETs url and returns the response body; a non-200 status wraps sentinel. +func httpGet(ctx context.Context, url string, sentinel error) (io.ReadCloser, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, http.NoBody) + if err != nil { + return nil, err + } + + resp, err := httpClient.Do(req) + if err != nil { + return nil, err + } + + if resp.StatusCode != http.StatusOK { + resp.Body.Close() + + return nil, fmt.Errorf("%w: %s", sentinel, resp.Status) + } + + return resp.Body, nil +} diff --git a/internal/usecase/packaging/archive_test.go b/internal/usecase/packaging/archive_test.go new file mode 100644 index 000000000..53abb66f7 --- /dev/null +++ b/internal/usecase/packaging/archive_test.go @@ -0,0 +1,148 @@ +package packaging + +import ( + "archive/tar" + "archive/zip" + "bytes" + "compress/gzip" + "errors" + "io" + "io/fs" + "testing" +) + +// neverendingReader is a Reader that fills any buffer with a constant byte and +// never returns io.EOF, used to drive readLimited past the cap without +// allocating a large buffer. +type neverendingReader struct{} + +func (neverendingReader) Read(p []byte) (int, error) { + for i := range p { + p[i] = 0xAB + } + + return len(p), nil +} + +func makeTarGz(t *testing.T, name string, content []byte) []byte { + t.Helper() + + var buf bytes.Buffer + + gz := gzip.NewWriter(&buf) + tw := tar.NewWriter(gz) + + hdr := &tar.Header{Name: name, Mode: 0o755, Size: int64(len(content))} + if err := tw.WriteHeader(hdr); err != nil { + t.Fatal(err) + } + + if _, err := tw.Write(content); err != nil { + t.Fatal(err) + } + + if err := tw.Close(); err != nil { + t.Fatal(err) + } + + if err := gz.Close(); err != nil { + t.Fatal(err) + } + + return buf.Bytes() +} + +func TestExtractBinaryReleaseTarGz(t *testing.T) { + t.Parallel() + + data := makeTarGz(t, "rpc_linux_x64", []byte("ELF-bytes")) + + name, content, err := extractBinary(data, "rpc_linux_x64.tar.gz") + if err != nil { + t.Fatal(err) + } + + if name != "rpc" || string(content) != "ELF-bytes" { + t.Fatalf("got (%q,%q)", name, content) + } +} + +func TestExtractBinaryBareExe(t *testing.T) { + t.Parallel() + + name, content, err := extractBinary([]byte("PE-bytes"), "rpc_windows_x64.exe") + if err != nil { + t.Fatal(err) + } + + if name != "rpc.exe" || string(content) != "PE-bytes" { + t.Fatalf("got (%q,%q)", name, content) + } +} + +func TestExtractBinaryUnexpectedEntry(t *testing.T) { + t.Parallel() + + data := makeTarGz(t, "README.md", []byte("x")) + + _, _, err := extractBinary(data, "rpc_linux_x64.tar.gz") + if !errors.Is(err, ErrBinaryNotFound) { + t.Fatalf("expected ErrBinaryNotFound, got %v", err) + } +} + +func TestBuildZip(t *testing.T) { + t.Parallel() + + out, err := buildZip([]zipEntry{{name: "rpc", content: []byte("bin")}}, []byte("cfg")) + if err != nil { + t.Fatal(err) + } + + zr, err := zip.NewReader(bytes.NewReader(out), int64(len(out))) + if err != nil { + t.Fatal(err) + } + + found := map[string]string{} + modes := map[string]fs.FileMode{} + + for _, f := range zr.File { + rc, err := f.Open() + if err != nil { + t.Fatal(err) + } + + var b bytes.Buffer + if _, err := b.ReadFrom(rc); err != nil { + t.Fatal(err) + } + + rc.Close() + + found[f.Name] = b.String() + modes[f.Name] = f.Mode() + } + + if found["rpc"] != "bin" || found["config.yaml"] != "cfg" { + t.Fatalf("unexpected zip contents: %v", found) + } + + if modes["rpc"]&0o111 == 0 { + t.Fatalf("expected rpc to be executable, mode = %v", modes["rpc"]) + } +} + +func TestReadLimitedEntryTooLarge(t *testing.T) { + t.Parallel() + + // Provide a reader that yields exactly maxArchiveBytes+10 bytes before being + // cut off by io.LimitReader — enough to exceed the cap without allocating the + // full buffer in memory. + r := io.LimitReader(neverendingReader{}, maxArchiveBytes+10) + + _, err := readLimited(r) + if !errors.Is(err, ErrEntryTooLarge) { + t.Fatalf("expected ErrEntryTooLarge, got: %v", err) + } +} diff --git a/internal/usecase/packaging/github.go b/internal/usecase/packaging/github.go new file mode 100644 index 000000000..f9f1aa498 --- /dev/null +++ b/internal/usecase/packaging/github.go @@ -0,0 +1,162 @@ +package packaging + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "path/filepath" + "regexp" + "slices" + "strconv" + "strings" + + "golang.org/x/mod/semver" + + dto "github.com/device-management-toolkit/console/internal/entity/dto/v1" + "github.com/device-management-toolkit/console/internal/entity/github" +) + +// assetRe matches rpc-go's published builds: rpc_linux_.tar.gz and rpc_windows_.exe. +var assetRe = regexp.MustCompile(`^rpc_(linux|windows)_([a-z0-9]+)\.(?:tar\.gz|exe)$`) + +// releasesURL builds the GitHub API releases URL for the given repo. +// base is overridable in tests (e.g. an httptest.Server URL). +func releasesURL(base, repo string) string { + return fmt.Sprintf("%s/repos/%s/releases?per_page=%d", base, repo, maxReleases) +} + +const minSupportedMajor = 3 + +// maxReleases caps the releases offered, so one unpaginated GitHub page covers them. +const maxReleases = 5 + +// parseAsset extracts the os ("linux"/"windows") and arch from an rpc-go release +// asset filename. ok is false for non-build assets. +func parseAsset(filename string) (goos, arch string, ok bool) { + m := assetRe.FindStringSubmatch(filename) + if m == nil { + return "", "", false + } + + return m[1], m[2], true +} + +// isV3OrAbove reports whether a release tag is semver major >= 3 (betas count). +func isV3OrAbove(tag string) bool { + t := strings.TrimPrefix(strings.TrimSpace(tag), "v") + + dot := strings.IndexByte(t, '.') + if dot < 0 { + return false + } + + major, err := strconv.Atoi(t[:dot]) + if err != nil { + return false + } + + return major >= minSupportedMajor +} + +// ErrFetchReleases indicates the GitHub releases request did not return 200. +var ErrFetchReleases = errors.New("failed to fetch releases") + +// getReleases GETs a GitHub releases list URL and returns the raw release slice. +func getReleases(ctx context.Context, url string) ([]github.Release, error) { + body, err := httpGet(ctx, url, ErrFetchReleases) + if err != nil { + return nil, err + } + defer body.Close() + + var releases []github.Release + if err := json.NewDecoder(io.LimitReader(body, maxArchiveBytes)).Decode(&releases); err != nil { + return nil, err + } + + return releases, nil +} + +// filterReleases keeps the newest maxReleases v3+ releases and maps them to the +// UI DTO shape. +func filterReleases(releases []github.Release) []dto.RPCRelease { + out := make([]dto.RPCRelease, 0, len(releases)) + + for i := range releases { + if len(out) == maxReleases { + break + } + + r := &releases[i] + + if !isV3OrAbove(r.TagName) { + continue + } + + assets := make([]dto.RPCAsset, 0, len(r.Assets)) + + for _, a := range r.Assets { + if assetOS, arch, ok := parseAsset(a.Name); ok { + assets = append(assets, dto.RPCAsset{OS: assetOS, Arch: arch}) + } + } + + out = append(out, dto.RPCRelease{Version: r.TagName, Assets: assets}) + } + + return out +} + +// listLocalReleases scans an offline directory laid out as // +// and returns the newest maxReleases v3+ versions. +func listLocalReleases(dir string) ([]dto.RPCRelease, error) { + entries, err := os.ReadDir(dir) + if err != nil { + return nil, err + } + + out := make([]dto.RPCRelease, 0, len(entries)) + + for _, e := range entries { + if !e.IsDir() || !isV3OrAbove(e.Name()) { + continue + } + + version := e.Name() + + files, err := os.ReadDir(filepath.Join(dir, version)) + if err != nil { + return nil, err + } + + assets := make([]dto.RPCAsset, 0, len(files)) + + for _, f := range files { + if assetOS, arch, ok := parseAsset(f.Name()); ok { + assets = append(assets, dto.RPCAsset{OS: assetOS, Arch: arch}) + } + } + + if len(assets) > 0 { + out = append(out, dto.RPCRelease{Version: version, Assets: assets}) + } + } + + slices.SortFunc(out, func(a, b dto.RPCRelease) int { + return semver.Compare(canonicalTag(b.Version), canonicalTag(a.Version)) + }) + + if len(out) > maxReleases { + out = out[:maxReleases] + } + + return out, nil +} + +// canonicalTag adds the "v" prefix semver.Compare requires. +func canonicalTag(tag string) string { + return "v" + strings.TrimPrefix(tag, "v") +} diff --git a/internal/usecase/packaging/github_test.go b/internal/usecase/packaging/github_test.go new file mode 100644 index 000000000..cb9e26cc4 --- /dev/null +++ b/internal/usecase/packaging/github_test.go @@ -0,0 +1,218 @@ +package packaging + +import ( + "context" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "slices" + "testing" + + "github.com/device-management-toolkit/console/internal/entity/github" +) + +func TestParseAsset(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + filename string + wantOK bool + wantOS string + wantArch string + }{ + {"linux", "rpc_linux_x64.tar.gz", true, "linux", "x64"}, + {"windows", "rpc_windows_x86.exe", true, "windows", "x86"}, + {"shared library skipped", "rpc_so_x64.tar.gz", false, "", ""}, + {"signature skipped", "rpc_windows_x64.exe.sigstore.json", false, "", ""}, + {"licenses skipped", "licenses.zip", false, "", ""}, + {"source skipped", "Source code (zip)", false, "", ""}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + goos, arch, ok := parseAsset(tc.filename) + if ok != tc.wantOK || goos != tc.wantOS || arch != tc.wantArch { + t.Fatalf("parseAsset(%q) = (%q,%q,%v), want (%q,%q,%v)", + tc.filename, goos, arch, ok, tc.wantOS, tc.wantArch, tc.wantOK) + } + }) + } +} + +func TestIsV3OrAbove(t *testing.T) { + t.Parallel() + + cases := map[string]bool{ + "v3.0.1": true, "v3.1.0-beta": true, "v4.0.0": true, + "v2.9.9": false, "v1.0.0": false, "not-a-tag": false, + } + for tag, want := range cases { + if got := isV3OrAbove(tag); got != want { + t.Fatalf("isV3OrAbove(%q) = %v, want %v", tag, got, want) + } + } +} + +func TestGetReleasesOnline(t *testing.T) { + t.Parallel() + + body := `[ + {"tag_name":"v3.0.1","prerelease":false,"assets":[ + {"name":"rpc_linux_x64.tar.gz","browser_download_url":"http://x/l"}, + {"name":"rpc_windows_x64.exe","browser_download_url":"http://x/w"}]}, + {"tag_name":"v2.9.0","prerelease":false,"assets":[ + {"name":"rpc_linux_x64.tar.gz","browser_download_url":"http://x/old"}]} + ]` + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte(body)) + })) + defer srv.Close() + + fetched, err := getReleases(context.Background(), srv.URL) + if err != nil { + t.Fatal(err) + } + + rels := filterReleases(fetched) + + if len(rels) != 1 || rels[0].Version != "v3.0.1" || len(rels[0].Assets) != 2 { + t.Fatalf("unexpected releases: %+v", rels) + } +} + +func TestListLocalReleases(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + + verDir := filepath.Join(dir, "v3.0.1") + + if err := os.MkdirAll(verDir, 0o750); err != nil { + t.Fatal(err) + } + + if err := os.WriteFile(filepath.Join(verDir, "rpc_linux_x64.tar.gz"), []byte("x"), 0o600); err != nil { + t.Fatal(err) + } + + rels, err := listLocalReleases(dir) + if err != nil { + t.Fatal(err) + } + + if len(rels) != 1 || rels[0].Version != "v3.0.1" || len(rels[0].Assets) != 1 || + rels[0].Assets[0].OS != "linux" || rels[0].Assets[0].Arch != "x64" { + t.Fatalf("unexpected local releases: %+v", rels) + } +} + +func TestGetReleasesHTTPError(t *testing.T) { + t.Parallel() + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + defer srv.Close() + + _, err := getReleases(context.Background(), srv.URL) + if err == nil { + t.Fatal("expected error on non-200 response, got nil") + } + + if !errors.Is(err, ErrFetchReleases) { + t.Fatalf("expected error to wrap ErrFetchReleases, got %v", err) + } +} + +// Only the newest few releases are offered, so the releases call stays a single +// unpaginated request. +func TestFilterReleasesCapsAtMaxReleases(t *testing.T) { + t.Parallel() + + releases := make([]github.Release, 0, 12) + for i := range 12 { + releases = append(releases, github.Release{ + TagName: fmt.Sprintf("v3.0.%d", i), + Assets: []github.Asset{{Name: "rpc_linux_x64.tar.gz"}}, + }) + } + + got := filterReleases(releases) + if len(got) != maxReleases { + t.Fatalf("filterReleases() returned %d releases, want %d", len(got), maxReleases) + } + + // The cap must keep the newest, which GitHub returns first. + if got[0].Version != "v3.0.0" { + t.Errorf("first release = %q, want the first returned by the API", got[0].Version) + } +} + +// Pre-v3 tags must not consume cap slots that a supported release could fill. +func TestFilterReleasesSkipsPreV3WithinCap(t *testing.T) { + t.Parallel() + + releases := []github.Release{ + {TagName: "v2.9.0", Assets: []github.Asset{{Name: "rpc_linux_x64.tar.gz"}}}, + {TagName: "v3.1.0", Assets: []github.Asset{{Name: "rpc_linux_x64.tar.gz"}}}, + {TagName: "v3.0.0", Assets: []github.Asset{{Name: "rpc_linux_x64.tar.gz"}}}, + } + + got := filterReleases(releases) + if len(got) != 2 { + t.Fatalf("filterReleases() returned %d releases, want 2", len(got)) + } + + if got[0].Version != "v3.1.0" || got[1].Version != "v3.0.0" { + t.Errorf("unexpected releases: %+v", got) + } +} + +func TestReleasesURLRequestsOnePage(t *testing.T) { + t.Parallel() + + got := releasesURL("https://api.github.com", "owner/repo") + want := "https://api.github.com/repos/owner/repo/releases?per_page=5" + + if got != want { + t.Errorf("releasesURL() = %q, want %q", got, want) + } +} + +func TestListLocalReleasesNewestFirstV3Only(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + + for _, v := range []string{"v2.9.0", "v3.0.0", "v3.10.0", "v3.2.0", "v3.1.0", "v3.3.0-beta.1", "v3.3.0", "v3.9.0"} { + verDir := filepath.Join(dir, v) + + if err := os.MkdirAll(verDir, 0o750); err != nil { + t.Fatal(err) + } + + if err := os.WriteFile(filepath.Join(verDir, "rpc_linux_x64.tar.gz"), []byte("x"), 0o600); err != nil { + t.Fatal(err) + } + } + + rels, err := listLocalReleases(dir) + if err != nil { + t.Fatal(err) + } + + got := make([]string, 0, len(rels)) + for _, r := range rels { + got = append(got, r.Version) + } + + if want := []string{"v3.10.0", "v3.9.0", "v3.3.0", "v3.3.0-beta.1", "v3.2.0"}; !slices.Equal(got, want) { + t.Fatalf("versions = %v, want %v", got, want) + } +} From 26c31fb78ec328beb5bc5dab62de141e654e2d9d Mon Sep 17 00:00:00 2001 From: Mike Date: Fri, 18 Sep 2026 16:26:26 -0700 Subject: [PATCH 2/3] refactor(api): add rpc-go packaging service Builds the Download RPC zip; not yet exposed over HTTP. - Resolve the requested build from GitHub, falling back to package.local_dir - package.disable_fetch serves builds from local_dir only; with fetching on, GitHub releases list first, then local-only versions - Render rpc-go's config.yaml for activate or deactivate with token, userpass, or no embedded credentials, scoped to the caller's tenant - Mint the auth token with a requested lifetime capped by package.max_token_ttl; no token is minted when auth is disabled - Point rpc-go at the request's serverUrl or the listener address, and skip cert checks when the listener serves a generated certificate - Package the Windows and Linux builds together for os "both" --- .env.example | 10 + config/config.go | 25 + config/config.yml | 3 + config/config_test.go | 44 ++ internal/usecase/packaging/archive.go | 13 + internal/usecase/packaging/config.go | 149 +++++ internal/usecase/packaging/config_test.go | 261 ++++++++ internal/usecase/packaging/github.go | 37 ++ internal/usecase/packaging/interface.go | 14 + internal/usecase/packaging/packaging.go | 289 +++++++++ internal/usecase/packaging/packaging_test.go | 638 +++++++++++++++++++ internal/usecase/packaging/scheme_test.go | 134 ++++ internal/usecase/packaging/token.go | 62 ++ internal/usecase/packaging/token_test.go | 106 +++ 14 files changed, 1785 insertions(+) create mode 100644 internal/usecase/packaging/config.go create mode 100644 internal/usecase/packaging/config_test.go create mode 100644 internal/usecase/packaging/interface.go create mode 100644 internal/usecase/packaging/packaging.go create mode 100644 internal/usecase/packaging/packaging_test.go create mode 100644 internal/usecase/packaging/scheme_test.go create mode 100644 internal/usecase/packaging/token.go create mode 100644 internal/usecase/packaging/token_test.go diff --git a/.env.example b/.env.example index 145883c6a..bc89c5009 100644 --- a/.env.example +++ b/.env.example @@ -82,3 +82,13 @@ GIN_MODE=release AUTH_CLIENT_ID= # ex. "https://login.microsoftonline.com//v2.0 for Azure Entra -- used for discovery AUTH_ISSUER= +# DOWNLOAD RPC PACKAGING +# GitHub repository the rpc-go releases are fetched from. +RPC_REPO=device-management-toolkit/rpc-go +# Ceiling for the auth-token lifetime a Download RPC package may request. +# Requests above it are rejected; 0 or unset uses 24h. +RPC_MAX_TOKEN_TTL=24h +# Directory of cached rpc-go builds, laid out as //. +RPC_LOCAL_DIR= +# true lists and serves builds from RPC_LOCAL_DIR only; false fetches from GitHub first, then RPC_LOCAL_DIR. +RPC_DISABLE_FETCH=false diff --git a/config/config.go b/config/config.go index 5bf5ab4ca..aa596db0d 100644 --- a/config/config.go +++ b/config/config.go @@ -26,6 +26,8 @@ var TrayMode bool var ( ErrJWTExpirationInvalid = errors.New("config: auth.jwtExpiration must be at least 1 minute (e.g. 24h) — very short expirations render tokens unusable") ErrRedirectionJWTExpirationInvalid = errors.New("config: auth.redirectionJWTExpiration must be at least 1 minute (e.g. 5m) — very short expirations render redirection tokens unusable") + ErrMaxTokenTTLInvalid = errors.New("config: package.max_token_ttl must be at least 1 minute (e.g. 24h), or 0 to use the default") + ErrLocalDirRequired = errors.New("config: package.local_dir is required when package.disable_fetch is true — set RPC_LOCAL_DIR or local_dir in config.yml") ErrJWTKeyMissing = errors.New("config: auth.jwtKey is required — set AUTH_JWT_KEY environment variable or jwtKey in config.yml to a strong secret") ) @@ -54,6 +56,7 @@ type ( EA `yaml:"ea"` Auth `yaml:"auth"` UI `yaml:"ui"` + Package `yaml:"package"` } // App -. @@ -152,6 +155,16 @@ type ( UI struct { ExternalURL string `yaml:"externalUrl" env:"UI_EXTERNAL_URL"` } + + // Package -. Settings for the Download RPC packaging endpoints. + Package struct { + RPCRepo string `yaml:"rpc_repo" env:"RPC_REPO"` + LocalDir string `yaml:"local_dir" env:"RPC_LOCAL_DIR"` + // DisableFetch serves rpc-go builds from LocalDir only, never contacting GitHub. + DisableFetch bool `yaml:"disable_fetch" env:"RPC_DISABLE_FETCH"` + // MaxTokenTTL caps the auth-token lifetime a package request may ask for. + MaxTokenTTL time.Duration `yaml:"max_token_ttl" env:"RPC_MAX_TOKEN_TTL"` + } ) // CookieAuthEnabled reports whether the HttpOnly session cookie is in use. Off @@ -251,6 +264,10 @@ func defaultConfig() *Config { UI: UI{ ExternalURL: "", }, + Package: Package{ + RPCRepo: "device-management-toolkit/rpc-go", + MaxTokenTTL: 24 * time.Hour, + }, } } @@ -471,6 +488,14 @@ func (c *Config) validate() error { return ErrRedirectionJWTExpirationInvalid } + if c.MaxTokenTTL != 0 && c.MaxTokenTTL < time.Minute { + return ErrMaxTokenTTLInvalid + } + + if c.DisableFetch && c.LocalDir == "" { + return ErrLocalDirRequired + } + if !c.Disabled && c.JWTKey == "" { return ErrJWTKeyMissing } diff --git a/config/config.yml b/config/config.yml index e88910151..120789e01 100644 --- a/config/config.yml +++ b/config/config.yml @@ -59,4 +59,7 @@ ui: # - Ignored: When building without 'noui' tag (embedded UI is served normally) # Example: https://ui.example.com externalUrl: "" +package: + rpc_repo: device-management-toolkit/rpc-go + local_dir: "" diff --git a/config/config_test.go b/config/config_test.go index f968b8696..39a854463 100644 --- a/config/config_test.go +++ b/config/config_test.go @@ -457,6 +457,50 @@ func TestValidate_SubMinuteJWTExpiration(t *testing.T) { require.ErrorIs(t, err, ErrJWTExpirationInvalid) } +func TestValidate_SubMinuteMaxTokenTTL(t *testing.T) { + t.Parallel() + + cfg := defaultConfig() + cfg.MaxTokenTTL = 30 * time.Second + + err := cfg.validate() + require.ErrorIs(t, err, ErrMaxTokenTTLInvalid) +} + +func TestValidate_NegativeMaxTokenTTL(t *testing.T) { + t.Parallel() + + cfg := defaultConfig() + cfg.MaxTokenTTL = -1 * time.Hour + + err := cfg.validate() + require.ErrorIs(t, err, ErrMaxTokenTTLInvalid) +} + +func TestValidate_UnsetMaxTokenTTLIsAllowed(t *testing.T) { + t.Parallel() + + cfg := defaultConfig() + cfg.MaxTokenTTL = 0 + cfg.JWTKey = "test-jwt-key" + + require.NoError(t, cfg.validate()) +} + +func TestValidate_DisableFetchRequiresLocalDir(t *testing.T) { + t.Parallel() + + cfg := defaultConfig() + cfg.JWTKey = "test-jwt-key" + cfg.DisableFetch = true + + require.ErrorIs(t, cfg.validate(), ErrLocalDirRequired) + + cfg.LocalDir = "/opt/rpc-go" + + require.NoError(t, cfg.validate()) +} + func TestValidate_ZeroRedirectionJWTExpiration(t *testing.T) { t.Parallel() diff --git a/internal/usecase/packaging/archive.go b/internal/usecase/packaging/archive.go index 4d893f12b..363e6580e 100644 --- a/internal/usecase/packaging/archive.go +++ b/internal/usecase/packaging/archive.go @@ -31,6 +31,8 @@ var ( ErrBinaryNotFound = errors.New("rpc binary not found in archive") // ErrEntryTooLarge indicates an archive entry exceeded the size cap. ErrEntryTooLarge = errors.New("archive entry exceeds size limit") + // ErrDownloadAsset indicates a non-200 response downloading an asset. + ErrDownloadAsset = errors.New("failed to download asset") ) // readLimited reads up to maxArchiveBytes from r, guarding against decompression bombs. @@ -150,3 +152,14 @@ func httpGet(ctx context.Context, url string, sentinel error) (io.ReadCloser, er return resp.Body, nil } + +// downloadAsset fetches an asset's bytes over HTTP. +func downloadAsset(ctx context.Context, url string) ([]byte, error) { + body, err := httpGet(ctx, url, ErrDownloadAsset) + if err != nil { + return nil, err + } + defer body.Close() + + return readLimited(body) +} diff --git a/internal/usecase/packaging/config.go b/internal/usecase/packaging/config.go new file mode 100644 index 000000000..1380c06dd --- /dev/null +++ b/internal/usecase/packaging/config.go @@ -0,0 +1,149 @@ +package packaging + +import ( + "fmt" + "net/url" + + "gopkg.in/yaml.v2" + + dto "github.com/device-management-toolkit/console/internal/entity/dto/v1" +) + +const ( + exportPathFmt = "%s/api/v1/admin/profiles/export/%s" + authModeToken = "token" + authModeUserPass = "userpass" + commandActivate = "activate" + commandDeactivate = "deactivate" + defaultLMSAddress = "localhost" + defaultLMSPort = "16992" + defaultLogLevel = "info" +) + +// configFile mirrors the full rpc-go config.yaml structure. +// The configure subtree (amtfeatures/wired/wireless/tls/etc.) is omitted for v1; +// rpc-go ignores unused sections, so this is safe — only populate what you use. +type configFile struct { + LogLevel string `yaml:"log-level"` + JSON bool `yaml:"json"` + Verbose bool `yaml:"verbose"` + SkipCertCheck bool `yaml:"skip-cert-check"` + LMSAddress string `yaml:"lmsaddress"` + LMSPort string `yaml:"lmsport"` + // rpc-go resolves config keys by flag name: --tenantid and --password. + TenantID string `yaml:"tenantid"` + AMTPassword string `yaml:"password"` + // Omitted when empty: rpc-go reads config after env, so an empty key would + // override AUTH_TOKEN / AUTH_USERNAME / AUTH_PASSWORD set on the device. + AuthToken string `yaml:"auth-token,omitempty"` + AuthUsername string `yaml:"auth-username,omitempty"` + AuthPassword string `yaml:"auth-password,omitempty"` + AuthEndpoint string `yaml:"auth-endpoint"` + DevicesEndpoint string `yaml:"devices-endpoint"` + AMTInfo amtInfoConfig `yaml:"amtinfo"` + Activate activateConfig `yaml:"activate"` + Deactivate deactivateConfig `yaml:"deactivate"` +} + +// amtInfoConfig maps the amtinfo sub-section of rpc-go config.yaml. +type amtInfoConfig struct { + Ver bool `yaml:"ver"` + All bool `yaml:"all"` + SKU bool `yaml:"sku"` + UUID bool `yaml:"uuid"` + Mode bool `yaml:"mode"` + DNS bool `yaml:"dns"` + Hostname bool `yaml:"hostname"` + LAN bool `yaml:"lan"` + RAS bool `yaml:"ras"` + OperationalState bool `yaml:"operationalState"` + UserCert bool `yaml:"userCert"` + Build bool `yaml:"bld"` + Sync bool `yaml:"sync"` + URL string `yaml:"url"` +} + +// activateConfig maps the activate sub-section of rpc-go config.yaml. +type activateConfig struct { + Local bool `yaml:"local"` + URL string `yaml:"url"` + Profile string `yaml:"profile"` + Proxy string `yaml:"proxy"` + CCM bool `yaml:"ccm"` + ACM bool `yaml:"acm"` + Key string `yaml:"key"` + DNS string `yaml:"dns"` + Hostname string `yaml:"hostname"` + Name string `yaml:"name"` + UUID string `yaml:"uuid"` + StopConfig bool `yaml:"stopConfig"` + SkipIPRenew bool `yaml:"skipIPRenew"` + ProvisioningCert string `yaml:"provisioningCert"` + ProvisioningCertPwd string `yaml:"provisioningCertPwd"` +} + +// deactivateConfig maps the deactivate sub-section of rpc-go config.yaml. +type deactivateConfig struct { + URL string `yaml:"url"` + Profile string `yaml:"profile"` + Proxy string `yaml:"proxy"` +} + +// defaultConfigFile returns a configFile pre-filled with rpc-go sample defaults. +func defaultConfigFile() configFile { + return configFile{ + LogLevel: defaultLogLevel, + LMSAddress: defaultLMSAddress, + LMSPort: defaultLMSPort, + } +} + +// configInputs carries caller-resolved values (e.g. pre-minted auth token) +// so that renderConfig remains a pure function. +type configInputs struct { + AuthEndpoint string + DevicesEndpoint string + ExportBase string // base URL for the activate profile-export URL + AuthToken string // non-empty when auth mode == token + TenantID string // tenant the package is scoped to; empty is the default tenant + SkipCertCheck bool // server presents a self-signed cert rpc-go cannot chain +} + +// renderConfig builds a complete rpc-go config.yaml from a PackageRequest and +// resolved configInputs. It starts from defaults and applies only the +// request-driven fields; everything else keeps zero/default values. +func renderConfig(req dto.PackageRequest, in configInputs) ([]byte, error) { + cfg := defaultConfigFile() + + cfg.AuthEndpoint = in.AuthEndpoint + cfg.DevicesEndpoint = in.DevicesEndpoint + cfg.TenantID = in.TenantID + cfg.SkipCertCheck = in.SkipCertCheck + + // Mode "none" writes no credentials; the operator supplies them on the device. + switch req.Auth.Mode { + case authModeToken: + cfg.AuthToken = in.AuthToken + case authModeUserPass: + cfg.AuthUsername = req.Auth.Username + cfg.AuthPassword = req.Auth.Password + } + + switch req.Command { + case commandActivate: + cfg.Activate.URL = fmt.Sprintf(exportPathFmt, in.ExportBase, url.PathEscape(req.Profile)) + if req.Domain != "" { + cfg.Activate.URL += "?domainName=" + url.QueryEscape(req.Domain) + } + case commandDeactivate: + // Remote deactivate targets the server's devices API; the shared auth block carries credentials. + cfg.Deactivate.URL = in.DevicesEndpoint + } + + out, err := yaml.Marshal(cfg) //nolint:gosec // G117: config intentionally serializes credential fields (password, auth-token, etc.) — the caller is responsible for protecting the resulting bytes + if err != nil { + return nil, fmt.Errorf("marshal config: %w", err) + } + + return out, nil +} diff --git a/internal/usecase/packaging/config_test.go b/internal/usecase/packaging/config_test.go new file mode 100644 index 000000000..54b6b7078 --- /dev/null +++ b/internal/usecase/packaging/config_test.go @@ -0,0 +1,261 @@ +package packaging + +import ( + "strings" + "testing" + + "gopkg.in/yaml.v2" + + dto "github.com/device-management-toolkit/console/internal/entity/dto/v1" +) + +func unmarshalConfig(t *testing.T, data []byte) map[string]interface{} { + t.Helper() + + var m map[string]interface{} + + if err := yaml.Unmarshal(data, &m); err != nil { + t.Fatalf("result is not valid YAML: %v\n---\n%s", err, data) + } + + return m +} + +func activateSection(t *testing.T, m map[string]interface{}) map[interface{}]interface{} { + t.Helper() + + activate, ok := m["activate"].(map[interface{}]interface{}) + if !ok { + t.Fatalf("activate section missing or wrong type: %T", m["activate"]) + } + + return activate +} + +func deactivateSection(t *testing.T, m map[string]interface{}) map[interface{}]interface{} { + t.Helper() + + deactivate, ok := m["deactivate"].(map[interface{}]interface{}) + if !ok { + t.Fatalf("deactivate section missing or wrong type: %T", m["deactivate"]) + } + + return deactivate +} + +func TestRenderConfigTokenActivateDomain(t *testing.T) { + t.Parallel() + + req := dto.PackageRequest{ + Command: "activate", + Version: "v3.0.1", + OS: "linux", + Arch: "x86_64", + Auth: dto.PackageAuth{Mode: "token"}, + Profile: "myProfile", + Domain: "corp.com", + } + in := configInputs{ + AuthEndpoint: "https://auth.example.com/token", + DevicesEndpoint: "https://mps.example.com/devices", + ExportBase: "https://rps.example.com", + AuthToken: "tok", + } + + data, err := renderConfig(req, in) + if err != nil { + t.Fatalf("renderConfig returned error: %v", err) + } + + m := unmarshalConfig(t, data) + + if got, _ := m["auth-token"].(string); got != "tok" { + t.Errorf("auth-token = %q, want %q", got, "tok") + } + + activate := activateSection(t, m) + + activateURL, _ := activate["url"].(string) + + if !strings.Contains(activateURL, "/profiles/export/myProfile") { + t.Errorf("activate.url = %q, want it to contain /profiles/export/myProfile", activateURL) + } + + if !strings.Contains(activateURL, "domainName=corp.com") { + t.Errorf("activate.url = %q, want it to contain domainName=corp.com", activateURL) + } +} + +func TestRenderConfigUserpassActivateNoDomain(t *testing.T) { + t.Parallel() + + req := dto.PackageRequest{ + Command: "activate", + Version: "v3.0.1", + OS: "linux", + Arch: "x86_64", + Auth: dto.PackageAuth{Mode: "userpass", Username: "admin", Password: "secret"}, + Profile: "p1", + Domain: "", + } + in := configInputs{ + AuthEndpoint: "https://auth.example.com/token", + DevicesEndpoint: "https://mps.example.com/devices", + ExportBase: "https://rps.example.com", + AuthToken: "", + } + + data, err := renderConfig(req, in) + if err != nil { + t.Fatalf("renderConfig returned error: %v", err) + } + + m := unmarshalConfig(t, data) + + if got, _ := m["auth-username"].(string); got != "admin" { + t.Errorf("auth-username = %q, want %q", got, "admin") + } + + if got, _ := m["auth-password"].(string); got != "secret" { + t.Errorf("auth-password = %q, want %q", got, "secret") + } + + activate := activateSection(t, m) + + activateURL, _ := activate["url"].(string) + + if !strings.Contains(activateURL, "/export/p1") { + t.Errorf("activate.url = %q, want it to contain /export/p1", activateURL) + } + + if strings.Contains(activateURL, "domainName") { + t.Errorf("activate.url = %q, should not contain domainName when domain is empty", activateURL) + } +} + +func TestRenderConfigTokenDeactivate(t *testing.T) { + t.Parallel() + + req := dto.PackageRequest{ + Command: "deactivate", + Version: "v3.0.1", + OS: "linux", + Arch: "x86_64", + Auth: dto.PackageAuth{Mode: "token"}, + } + in := configInputs{ + AuthEndpoint: "https://auth.example.com/token", + DevicesEndpoint: "https://mps.example.com/devices", + ExportBase: "https://rps.example.com", + AuthToken: "tok", + } + + data, err := renderConfig(req, in) + if err != nil { + t.Fatalf("renderConfig returned error: %v", err) + } + + m := unmarshalConfig(t, data) + + if got, _ := m["auth-token"].(string); got != "tok" { + t.Errorf("auth-token = %q, want %q", got, "tok") + } + + if activate, ok := m["activate"].(map[interface{}]interface{}); ok { + if activateURL, _ := activate["url"].(string); activateURL != "" { + t.Errorf("activate.url = %q, want empty for deactivate command", activateURL) + } + } + + deactivate := deactivateSection(t, m) + + deactivateURL, _ := deactivate["url"].(string) + + if deactivateURL == "" { + t.Errorf("deactivate.url is empty, want non-empty") + } +} + +// Mode "none" must leave the credential keys out entirely, not write them empty: +// rpc-go applies config.yaml after env, so an empty key would override the +// AUTH_* variables the operator sets on the device. +func TestRenderConfigNoneOmitsCredentialKeys(t *testing.T) { + t.Parallel() + + req := dto.PackageRequest{ + Command: "activate", + Auth: dto.PackageAuth{Mode: "none"}, + Profile: "p1", + } + + out, err := renderConfig(req, configInputs{ + AuthEndpoint: "https://console.example/api/v1/authorize", + DevicesEndpoint: "https://console.example/api/v1/devices", + ExportBase: "https://console.example", + AuthToken: "must-not-appear", + }) + if err != nil { + t.Fatal(err) + } + + m := unmarshalConfig(t, out) + + for _, key := range []string{"auth-token", "auth-username", "auth-password"} { + if _, present := m[key]; present { + t.Errorf("%s present in config for mode none; it must be omitted", key) + } + } + + // rpc-go still needs to know where to exchange the credentials it is given. + if got, _ := m["auth-endpoint"].(string); got != "https://console.example/api/v1/authorize" { + t.Errorf("auth-endpoint = %q, want it kept for mode none", got) + } +} + +// Token mode must not leak username/password keys, even empty ones. +func TestRenderConfigTokenOmitsUserPassKeys(t *testing.T) { + t.Parallel() + + req := dto.PackageRequest{ + Command: "deactivate", + Auth: dto.PackageAuth{Mode: "token"}, + } + + out, err := renderConfig(req, configInputs{AuthToken: "tok"}) + if err != nil { + t.Fatal(err) + } + + m := unmarshalConfig(t, out) + + if got, _ := m["auth-token"].(string); got != "tok" { + t.Errorf("auth-token = %q, want %q", got, "tok") + } + + for _, key := range []string{"auth-username", "auth-password"} { + if _, present := m[key]; present { + t.Errorf("%s present in config for token mode", key) + } + } +} + +func TestRenderConfigEscapesProfileName(t *testing.T) { + t.Parallel() + + req := dto.PackageRequest{ + Command: "activate", + Auth: dto.PackageAuth{Mode: "none"}, + Profile: "My Profile", + Domain: "corp", + } + + out, err := renderConfig(req, configInputs{ExportBase: "https://console.example.com"}) + if err != nil { + t.Fatal(err) + } + + got := activateSection(t, unmarshalConfig(t, out))["url"] + if want := "https://console.example.com/api/v1/admin/profiles/export/My%20Profile?domainName=corp"; got != want { + t.Errorf("activate url = %v, want %v", got, want) + } +} diff --git a/internal/usecase/packaging/github.go b/internal/usecase/packaging/github.go index f9f1aa498..0ca01b406 100644 --- a/internal/usecase/packaging/github.go +++ b/internal/usecase/packaging/github.go @@ -80,6 +80,24 @@ func getReleases(ctx context.Context, url string) ([]github.Release, error) { return releases, nil } +// findAsset searches releases for an asset matching version, goos, and arch. +// It returns the download URL, asset name, and whether a match was found. +func findAsset(releases []github.Release, version, goos, arch string) (url, name string, ok bool) { + for i := range releases { + if releases[i].TagName != version { + continue + } + + for _, a := range releases[i].Assets { + if aos, aarch, parsed := parseAsset(a.Name); parsed && aos == goos && aarch == arch { + return a.BrowserDownloadURL, a.Name, true + } + } + } + + return "", "", false +} + // filterReleases keeps the newest maxReleases v3+ releases and maps them to the // UI DTO shape. func filterReleases(releases []github.Release) []dto.RPCRelease { @@ -110,6 +128,25 @@ func filterReleases(releases []github.Release) []dto.RPCRelease { return out } +// mergeReleases appends local versions missing from fetched; fetched wins on a +// shared version. +func mergeReleases(fetched, local []dto.RPCRelease) []dto.RPCRelease { + seen := make(map[string]bool, len(fetched)) + out := append(make([]dto.RPCRelease, 0, len(fetched)+len(local)), fetched...) + + for _, r := range fetched { + seen[r.Version] = true + } + + for _, r := range local { + if !seen[r.Version] { + out = append(out, r) + } + } + + return out +} + // listLocalReleases scans an offline directory laid out as // // and returns the newest maxReleases v3+ versions. func listLocalReleases(dir string) ([]dto.RPCRelease, error) { diff --git a/internal/usecase/packaging/interface.go b/internal/usecase/packaging/interface.go new file mode 100644 index 000000000..0f20607b2 --- /dev/null +++ b/internal/usecase/packaging/interface.go @@ -0,0 +1,14 @@ +package packaging + +import ( + "context" + "io" + + dto "github.com/device-management-toolkit/console/internal/entity/dto/v1" +) + +// Feature is the public contract for the packaging usecase. +type Feature interface { + ListVersions(ctx context.Context) ([]dto.RPCRelease, error) + BuildPackage(ctx context.Context, req dto.PackageRequest, tenantID string) (io.Reader, string, error) +} diff --git a/internal/usecase/packaging/packaging.go b/internal/usecase/packaging/packaging.go new file mode 100644 index 000000000..3eb3587b7 --- /dev/null +++ b/internal/usecase/packaging/packaging.go @@ -0,0 +1,289 @@ +package packaging + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "os" + "path/filepath" + "strings" + + "github.com/device-management-toolkit/console/config" + dto "github.com/device-management-toolkit/console/internal/entity/dto/v1" + "github.com/device-management-toolkit/console/internal/entity/github" + "github.com/device-management-toolkit/console/pkg/logger" +) + +const ( + githubDefaultBase = "https://api.github.com" + + schemeHTTP = "http" + schemeHTTPS = "https" + + // osBoth requests the Windows and Linux builds in one package. + osBoth = "both" +) + +// listenerScheme returns the URL scheme the HTTP listener serves. +func listenerScheme(tlsEnabled bool) string { + if tlsEnabled { + return schemeHTTPS + } + + return schemeHTTP +} + +// ErrAssetNotFound is returned when the requested asset cannot be found online or locally. +var ErrAssetNotFound = errors.New("asset not found") + +// ErrUnsafeVersion is returned when req.Version contains path-traversal characters. +var ErrUnsafeVersion = errors.New("unsafe version: contains path separator or dot-dot") + +// Service implements the Feature interface for building rpc-go download packages. +type Service struct { + cfg *config.Config + l logger.Interface + githubBase string +} + +// New constructs a Service with the default GitHub API base URL. +func New(cfg *config.Config, l logger.Interface) *Service { + return &Service{ + cfg: cfg, + l: l, + githubBase: githubDefaultBase, + } +} + +// ListVersions returns the available rpc-go releases: GitHub's first, then any +// local-only versions. With fetching disabled only the local directory is read; +// if GitHub fails the local directory is the fallback. +func (s *Service) ListVersions(ctx context.Context) ([]dto.RPCRelease, error) { + if s.cfg.DisableFetch { + return listLocalReleases(s.cfg.LocalDir) + } + + fetched, err := getReleases(ctx, releasesURL(s.githubBase, s.cfg.RPCRepo)) + if err != nil { + if s.cfg.LocalDir == "" { + return nil, err + } + + s.l.Warn("github fetch failed, falling back to local dir: %v", err) + + return listLocalReleases(s.cfg.LocalDir) + } + + releases := filterReleases(fetched) + + if s.cfg.LocalDir == "" { + return releases, nil + } + + local, err := listLocalReleases(s.cfg.LocalDir) + if err != nil { + s.l.Warn("read local dir, listing github releases only: %v", err) + + return releases, nil + } + + return mergeReleases(releases, local), nil +} + +// BuildPackage resolves the requested rpc-go binary, renders a config.yaml, and +// returns a zip reader together with a suggested filename. tenantID scopes the +// generated config to the caller's tenant. +func (s *Service) BuildPackage(ctx context.Context, req dto.PackageRequest, tenantID string) (io.Reader, string, error) { + oses := packageOSes(req.OS) + binaries := make([]zipEntry, 0, len(oses)) + + var ( + releases []github.Release + releaseErr error + ) + + if !s.cfg.DisableFetch { + releases, releaseErr = getReleases(ctx, releasesURL(s.githubBase, s.cfg.RPCRepo)) + } + + for _, goos := range oses { + data, assetName, err := s.resolveAsset(ctx, releases, releaseErr, req.Version, goos, req.Arch) + if err != nil { + return nil, "", fmt.Errorf("resolve asset: %w", err) + } + + binName, binary, err := extractBinary(data, assetName) + if err != nil { + return nil, "", fmt.Errorf("extract binary: %w", err) + } + + binaries = append(binaries, zipEntry{name: binName, content: binary}) + } + + inputs, err := s.buildConfigInputs(req, tenantID) + if err != nil { + return nil, "", fmt.Errorf("build config inputs: %w", err) + } + + cfgYAML, err := renderConfig(req, inputs) + if err != nil { + return nil, "", fmt.Errorf("render config: %w", err) + } + + zipBytes, err := buildZip(binaries, cfgYAML) + if err != nil { + return nil, "", fmt.Errorf("build zip: %w", err) + } + + filename := fmt.Sprintf("rpc-%s-%s-%s.zip", safeFilenamePart(req.Command), safeFilenamePart(req.OS), safeFilenamePart(req.Arch)) + + return bytes.NewReader(zipBytes), filename, nil +} + +// fetchOnline downloads the matching asset from the fetched GitHub releases. +func fetchOnline(ctx context.Context, releases []github.Release, version, goos, arch string) (data []byte, assetName string, err error) { + assetURL, name, found := findAsset(releases, version, goos, arch) + if !found { + return nil, "", fmt.Errorf("%w: version=%s os=%s arch=%s", ErrAssetNotFound, version, goos, arch) + } + + data, err = downloadAsset(ctx, assetURL) + if err != nil { + return nil, "", err + } + + return data, name, nil +} + +// packageOSes lists the builds a request packages; "both" means Windows and Linux. +func packageOSes(goos string) []string { + if goos == osBoth { + return []string{"windows", "linux"} + } + + return []string{goos} +} + +// resolveAsset downloads the asset from the fetched releases, falling back to the +// local directory when configured; with fetching disabled only the local directory is read. +func (s *Service) resolveAsset(ctx context.Context, releases []github.Release, releaseErr error, version, goos, arch string) (data []byte, assetName string, err error) { + if s.cfg.DisableFetch { + return findLocalAsset(s.cfg.LocalDir, version, goos, arch) + } + + onlineErr := releaseErr + if onlineErr == nil { + data, assetName, onlineErr = fetchOnline(ctx, releases, version, goos, arch) + if onlineErr == nil { + return data, assetName, nil + } + } + + if s.cfg.LocalDir != "" { + s.l.Warn("asset not available online, trying local dir: %v", onlineErr) + + return findLocalAsset(s.cfg.LocalDir, version, goos, arch) + } + + return nil, "", onlineErr +} + +// safeFilenamePart replaces characters unsafe in a Content-Disposition filename with a hyphen. +func safeFilenamePart(s string) string { + return strings.Map(func(r rune) rune { + switch { + case r >= 'a' && r <= 'z', r >= 'A' && r <= 'Z', r >= '0' && r <= '9', r == '.', r == '-', r == '_': + return r + default: + return '-' + } + }, s) +} + +// validateVersion rejects a version that is not a single safe path element under LocalDir. +func validateVersion(version string) error { + if version == "." || strings.ContainsAny(version, `/\`) || strings.Contains(version, "..") || version != filepath.Base(version) { + return fmt.Errorf("%w: %q", ErrUnsafeVersion, version) + } + + return nil +} + +// findLocalAsset returns the matching asset file under //. +func findLocalAsset(dir, version, goos, arch string) (data []byte, assetName string, err error) { + if err := validateVersion(version); err != nil { + return nil, "", err + } + + versionDir := filepath.Join(dir, version) + + entries, rdErr := os.ReadDir(versionDir) + if rdErr != nil { + return nil, "", fmt.Errorf("read local version dir: %w", rdErr) + } + + for _, e := range entries { + if assetOS, assetArch, ok := parseAsset(e.Name()); ok && assetOS == goos && assetArch == arch { + b, readErr := os.ReadFile(filepath.Join(versionDir, e.Name())) + if readErr != nil { + return nil, "", fmt.Errorf("read local asset: %w", readErr) + } + + return b, e.Name(), nil + } + } + + return nil, "", fmt.Errorf("%w: version=%s os=%s arch=%s", ErrAssetNotFound, version, goos, arch) +} + +// buildConfigInputs resolves the server URLs and, for token auth, mints the JWT +// that renderConfig writes into config.yaml. +func (s *Service) buildConfigInputs(req dto.PackageRequest, tenantID string) (configInputs, error) { + base := strings.TrimRight(req.ServerURL, "/") + + if base == "" { + // Best-effort default for callers that do not supply a server URL. + host := s.cfg.Host + if host == "" { + host = "localhost" + } + + base = fmt.Sprintf("%s://%s:%s", listenerScheme(s.cfg.TLS.Enabled), host, s.cfg.Port) + } + + // Without a configured cert the listener serves a generated self-signed one, + // which rpc-go cannot chain to a trusted root. + skipCertCheck := s.cfg.TLS.Enabled && s.cfg.TLS.CertFile == "" + + in := configInputs{ + AuthEndpoint: base + "/api/v1/authorize", + DevicesEndpoint: base + "/api/v1/devices", + ExportBase: base, + TenantID: tenantID, + SkipCertCheck: skipCertCheck, + } + + // With auth disabled the server accepts requests without a token, so none is minted. + if req.Auth.Mode == authModeToken && !s.cfg.Disabled { + // The OIDC verifier rejects Console-signed tokens. + if s.cfg.ClientID != "" { + return configInputs{}, ErrTokenModeUnsupported + } + + ttl, ttlErr := resolveTokenTTL(req.TokenTTL, s.cfg.MaxTokenTTL) + if ttlErr != nil { + return configInputs{}, ttlErr + } + + tok, mintErr := mintToken(s.cfg.JWTKey, ttl) + if mintErr != nil { + return configInputs{}, fmt.Errorf("mint token: %w", mintErr) + } + + in.AuthToken = tok + } + + return in, nil +} diff --git a/internal/usecase/packaging/packaging_test.go b/internal/usecase/packaging/packaging_test.go new file mode 100644 index 000000000..6dd0ce80c --- /dev/null +++ b/internal/usecase/packaging/packaging_test.go @@ -0,0 +1,638 @@ +package packaging + +import ( + "archive/zip" + "bytes" + "context" + "errors" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "slices" + "testing" + "time" + + "github.com/golang-jwt/jwt/v5" + + "github.com/device-management-toolkit/console/config" + dto "github.com/device-management-toolkit/console/internal/entity/dto/v1" + "github.com/device-management-toolkit/console/pkg/logger" +) + +// newFailingServer returns an httptest.Server that always responds with 500 and +// registers t.Cleanup to close it. +func newFailingServer(t *testing.T) *httptest.Server { + t.Helper() + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + + t.Cleanup(srv.Close) + + return srv +} + +// newTestConfig returns a minimal *config.Config suitable for packaging tests. +func newTestConfig(localDir string) *config.Config { + return &config.Config{ + HTTP: config.HTTP{ + Host: "localhost", + Port: "8181", + }, + Auth: config.Auth{ + JWTKey: "test-jwt-key", + }, + Package: config.Package{ + RPCRepo: "device-management-toolkit/rpc-go", + LocalDir: localDir, + }, + } +} + +// buildOfflineFixture writes a real tar.gz containing an "rpc" binary to +// /v3.0.1/rpc_linux_x64.tar.gz and returns the tmp dir. +func buildOfflineFixture(t *testing.T) string { + t.Helper() + + tmp := t.TempDir() + verDir := filepath.Join(tmp, "v3.0.1") + + if err := os.MkdirAll(verDir, 0o750); err != nil { + t.Fatal(err) + } + + tarGzData := makeTarGz(t, "rpc_linux_x64", []byte("ELF-placeholder")) + + assetPath := filepath.Join(verDir, "rpc_linux_x64.tar.gz") + if err := os.WriteFile(assetPath, tarGzData, 0o600); err != nil { + t.Fatal(err) + } + + return tmp +} + +// newOfflineService constructs a Service backed by an httptest server that always +// returns 500 (forcing the online path to fail) and a local fixture directory. +func newOfflineService(t *testing.T, tmp string) *Service { + t.Helper() + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + + t.Cleanup(srv.Close) + + cfg := newTestConfig(tmp) + svc := New(cfg, logger.New("error")) + svc.githubBase = srv.URL + + return svc +} + +func TestListVersionsLocalFallback(t *testing.T) { + t.Parallel() + + tmp := buildOfflineFixture(t) + svc := newOfflineService(t, tmp) + + releases, err := svc.ListVersions(context.Background()) + if err != nil { + t.Fatalf("ListVersions returned error: %v", err) + } + + if len(releases) != 1 { + t.Fatalf("expected 1 release, got %d: %+v", len(releases), releases) + } + + if releases[0].Version != "v3.0.1" { + t.Errorf("release version = %q, want %q", releases[0].Version, "v3.0.1") + } + + if len(releases[0].Assets) != 1 { + t.Fatalf("expected 1 asset, got %d", len(releases[0].Assets)) + } + + if releases[0].Assets[0].OS != "linux" || releases[0].Assets[0].Arch != "x64" { + t.Errorf("asset = {OS:%q, Arch:%q}, want {OS:\"linux\", Arch:\"x64\"}", + releases[0].Assets[0].OS, releases[0].Assets[0].Arch) + } +} + +func TestBuildPackageDeactivateOffline(t *testing.T) { + t.Parallel() + + tmp := buildOfflineFixture(t) + svc := newOfflineService(t, tmp) + + req := dto.PackageRequest{ + Command: "deactivate", + Version: "v3.0.1", + OS: "linux", + Arch: "x64", + Auth: dto.PackageAuth{Mode: "token"}, + } + + reader, filename, err := svc.BuildPackage(context.Background(), req, "") + if err != nil { + t.Fatalf("BuildPackage returned error: %v", err) + } + + const wantFilename = "rpc-deactivate-linux-x64.zip" + if filename != wantFilename { + t.Errorf("filename = %q, want %q", filename, wantFilename) + } + + zipBytes, err := io.ReadAll(reader) + if err != nil { + t.Fatalf("reading zip bytes: %v", err) + } + + zr, err := zip.NewReader(bytes.NewReader(zipBytes), int64(len(zipBytes))) + if err != nil { + t.Fatalf("opening zip: %v", err) + } + + names := make(map[string]bool, len(zr.File)) + for _, f := range zr.File { + names[f.Name] = true + } + + if !names["rpc"] { + t.Errorf("zip does not contain 'rpc'; entries: %v", names) + } + + if !names["config.yaml"] { + t.Errorf("zip does not contain 'config.yaml'; entries: %v", names) + } +} + +// "both" packages the Windows and Linux builds beside one shared config. +func TestBuildPackageBothOffline(t *testing.T) { + t.Parallel() + + tmp := buildOfflineFixture(t) + + winPath := filepath.Join(tmp, "v3.0.1", "rpc_windows_x64.exe") + if err := os.WriteFile(winPath, []byte("PE-placeholder"), 0o600); err != nil { + t.Fatal(err) + } + + svc := newOfflineService(t, tmp) + + reader, _, err := svc.BuildPackage(context.Background(), dto.PackageRequest{ + Command: "deactivate", + Version: "v3.0.1", + OS: "both", + Arch: "x64", + Auth: dto.PackageAuth{Mode: "none"}, + }, "") + if err != nil { + t.Fatal(err) + } + + data, err := io.ReadAll(reader) + if err != nil { + t.Fatal(err) + } + + zr, err := zip.NewReader(bytes.NewReader(data), int64(len(data))) + if err != nil { + t.Fatal(err) + } + + names := map[string]bool{} + for _, f := range zr.File { + names[f.Name] = true + } + + for _, want := range []string{"rpc", "rpc.exe", "config.yaml"} { + if !names[want] { + t.Errorf("zip missing %s, has %v", want, names) + } + } + + if len(zr.File) != 3 { + t.Errorf("zip has %d entries, want 3", len(zr.File)) + } +} + +func TestBuildConfigInputsBaseURL(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + serverURL string + wantBase string + }{ + {"request server url wins", "https://override.example:8181", "https://override.example:8181"}, + {"trailing slash trimmed", "https://override.example:8181/", "https://override.example:8181"}, + {"falls back to listen address", "", "http://localhost:8181"}, + } + + for _, tt := range tests { + tt := tt + + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + svc := New(newTestConfig(""), logger.New("error")) + + in, err := svc.buildConfigInputs(dto.PackageRequest{ + Command: "activate", + Auth: dto.PackageAuth{Mode: "userpass"}, + ServerURL: tt.serverURL, + }, "") + if err != nil { + t.Fatal(err) + } + + if in.ExportBase != tt.wantBase { + t.Errorf("ExportBase = %q, want %q", in.ExportBase, tt.wantBase) + } + + if in.AuthEndpoint != tt.wantBase+"/api/v1/authorize" { + t.Errorf("AuthEndpoint = %q, want base %q", in.AuthEndpoint, tt.wantBase) + } + + if in.DevicesEndpoint != tt.wantBase+"/api/v1/devices" { + t.Errorf("DevicesEndpoint = %q, want base %q", in.DevicesEndpoint, tt.wantBase) + } + }) + } +} + +func TestBuildConfigInputsTokenTTL(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + requested string + maxTTL time.Duration + wantTTL time.Duration + wantErr error + }{ + {"defaults to an hour", "", 0, defaultTokenTTL, nil}, + {"honors the requested lifetime", "15m", 24 * time.Hour, 15 * time.Minute, nil}, + {"rejects a lifetime above the configured cap", "24h", time.Hour, 0, ErrTokenTTLTooLong}, + } + + for _, tt := range tests { + tt := tt + + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + cfg := newTestConfig("") + cfg.MaxTokenTTL = tt.maxTTL + + svc := New(cfg, logger.New("error")) + + in, err := svc.buildConfigInputs(dto.PackageRequest{ + Command: "activate", + Auth: dto.PackageAuth{Mode: "token"}, + TokenTTL: tt.requested, + }, "") + if !errors.Is(err, tt.wantErr) { + t.Fatalf("error = %v, want %v", err, tt.wantErr) + } + + if tt.wantErr != nil { + return + } + + claims := &jwt.RegisteredClaims{} + if _, perr := jwt.ParseWithClaims(in.AuthToken, claims, func(_ *jwt.Token) (interface{}, error) { + return []byte(cfg.JWTKey), nil + }); perr != nil { + t.Fatal(perr) + } + + if got := time.Until(claims.ExpiresAt.Time); got < tt.wantTTL-time.Minute || got > tt.wantTTL+time.Minute { + t.Errorf("token expires in %v, want roughly %v", got, tt.wantTTL) + } + }) + } +} + +func TestBuildPackagePathTraversalRejected(t *testing.T) { + t.Parallel() + + tmp := buildOfflineFixture(t) + svc := newOfflineService(t, tmp) + + req := dto.PackageRequest{ + Command: "deactivate", + Version: "../evil", + OS: "linux", + Arch: "x64", + Auth: dto.PackageAuth{Mode: "token"}, + } + + _, _, err := svc.BuildPackage(context.Background(), req, "") + if err == nil { + t.Fatal("expected error for path-traversal version, got nil") + } + + if !errors.Is(err, ErrUnsafeVersion) { + t.Errorf("expected ErrUnsafeVersion, got: %v", err) + } +} + +func TestValidateVersion(t *testing.T) { + t.Parallel() + + tests := []struct { + version string + wantErr bool + }{ + {"v3.0.1", false}, + {".", true}, + {"..", true}, + {"../x", true}, + {"a/b", true}, + {"a\\b", true}, + {"", true}, + } + + for _, tc := range tests { + t.Run(tc.version, func(t *testing.T) { + t.Parallel() + + err := validateVersion(tc.version) + if tc.wantErr && err == nil { + t.Fatalf("validateVersion(%q) = nil, want non-nil error", tc.version) + } + + if !tc.wantErr && err != nil { + t.Fatalf("validateVersion(%q) = %v, want nil", tc.version, err) + } + }) + } +} + +func TestListVersionsGitHubFailNoLocalDir(t *testing.T) { + t.Parallel() + + srv := newFailingServer(t) + + cfg := newTestConfig("") + svc := New(cfg, logger.New("error")) + svc.githubBase = srv.URL + + _, err := svc.ListVersions(context.Background()) + if err == nil { + t.Fatal("expected error when GitHub returns 500 and no LocalDir, got nil") + } + + if !errors.Is(err, ErrFetchReleases) { + t.Errorf("expected error to wrap ErrFetchReleases, got: %v", err) + } +} + +func TestSafeFilenamePart(t *testing.T) { + t.Parallel() + + tests := []struct { + input string + want string + }{ + {`linux"evil`, "linux-evil"}, + {"linux/etc/passwd", "linux-etc-passwd"}, + {"win\\path", "win-path"}, + {"v3.0.1", "v3.0.1"}, + {"x86_64", "x86_64"}, + {"activate", "activate"}, + {"hello world", "hello-world"}, + } + + for _, tc := range tests { + t.Run(tc.input, func(t *testing.T) { + t.Parallel() + + got := safeFilenamePart(tc.input) + if got != tc.want { + t.Errorf("safeFilenamePart(%q) = %q, want %q", tc.input, got, tc.want) + } + + for _, ch := range got { + if ch == '"' { + t.Errorf("safeFilenamePart(%q) = %q still contains double-quote", tc.input, got) + } + } + }) + } +} + +func TestBuildPackageOnline(t *testing.T) { + t.Parallel() + + tgz := makeTarGz(t, "rpc_linux_x64", []byte("ELF")) + + mux := http.NewServeMux() + + mux.HandleFunc("/dl/rpc.tar.gz", func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write(tgz) + }) + + // srvURL is set after the server is created; the closure captures the pointer. + var srvURL string + + mux.HandleFunc("/", func(w http.ResponseWriter, _ *http.Request) { + body := `[{"tag_name":"v3.0.1","assets":[{"name":"rpc_linux_x64.tar.gz","browser_download_url":"` + srvURL + `/dl/rpc.tar.gz"}]}]` + _, _ = w.Write([]byte(body)) + }) + + srv := httptest.NewServer(mux) + defer srv.Close() + + srvURL = srv.URL + + cfg := newTestConfig("") + cfg.RPCRepo = "owner/repo" + svc := New(cfg, logger.New("error")) + svc.githubBase = srv.URL + + reader, filename, err := svc.BuildPackage(context.Background(), dto.PackageRequest{ + Command: "activate", + Version: "v3.0.1", + OS: "linux", + Arch: "x64", + Auth: dto.PackageAuth{Mode: "token"}, + Profile: "p1", + }, "") + if err != nil { + t.Fatal(err) + } + + const wantFilename = "rpc-activate-linux-x64.zip" + if filename != wantFilename { + t.Fatalf("filename = %q, want %q", filename, wantFilename) + } + + zipBytes, err := io.ReadAll(reader) + if err != nil { + t.Fatalf("reading zip bytes: %v", err) + } + + zr, err := zip.NewReader(bytes.NewReader(zipBytes), int64(len(zipBytes))) + if err != nil { + t.Fatalf("opening zip: %v", err) + } + + names := make(map[string]bool, len(zr.File)) + for _, f := range zr.File { + names[f.Name] = true + } + + if !names["rpc"] { + t.Errorf("zip does not contain 'rpc'; entries: %v", names) + } + + if !names["config.yaml"] { + t.Errorf("zip does not contain 'config.yaml'; entries: %v", names) + } +} + +func TestBuildConfigInputsTokenOmittedWhenAuthDisabled(t *testing.T) { + t.Parallel() + + cfg := newTestConfig("") + cfg.Disabled = true + + svc := New(cfg, logger.New("error")) + + in, err := svc.buildConfigInputs(dto.PackageRequest{ + Command: "activate", + Auth: dto.PackageAuth{Mode: "token"}, + }, "") + if err != nil { + t.Fatal(err) + } + + if in.AuthToken != "" { + t.Errorf("AuthToken = %q, want empty when auth is disabled", in.AuthToken) + } +} + +func TestBuildConfigInputsTokenRejectedUnderOIDC(t *testing.T) { + t.Parallel() + + cfg := newTestConfig("") + cfg.ClientID = "console-client" + + svc := New(cfg, logger.New("error")) + + _, err := svc.buildConfigInputs(dto.PackageRequest{ + Command: "activate", + Auth: dto.PackageAuth{Mode: "token"}, + }, "") + if !errors.Is(err, ErrTokenModeUnsupported) { + t.Fatalf("err = %v, want ErrTokenModeUnsupported", err) + } +} + +// newReleasesServer serves body as the GitHub releases list and fails the test +// if called when fetching is disabled. +func newReleasesServer(t *testing.T, body string, allowed bool) *httptest.Server { + t.Helper() + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + if !allowed { + t.Errorf("github contacted with fetching disabled") + } + + _, _ = w.Write([]byte(body)) + })) + + t.Cleanup(srv.Close) + + return srv +} + +func TestListVersionsFetchedThenLocal(t *testing.T) { + t.Parallel() + + tmp := buildOfflineFixture(t) // v3.0.1 locally + + body := `[ + {"tag_name":"v3.1.0","assets":[{"name":"rpc_linux_x64.tar.gz","browser_download_url":"http://x/a"}]}, + {"tag_name":"v3.0.1","assets":[{"name":"rpc_windows_x64.exe","browser_download_url":"http://x/b"}]} + ]` + + if err := os.MkdirAll(filepath.Join(tmp, "v3.0.0"), 0o750); err != nil { + t.Fatal(err) + } + + if err := os.WriteFile(filepath.Join(tmp, "v3.0.0", "rpc_linux_x64.tar.gz"), []byte("x"), 0o600); err != nil { + t.Fatal(err) + } + + svc := New(newTestConfig(tmp), logger.New("error")) + svc.githubBase = newReleasesServer(t, body, true).URL + + releases, err := svc.ListVersions(context.Background()) + if err != nil { + t.Fatal(err) + } + + got := make([]string, 0, len(releases)) + for _, r := range releases { + got = append(got, r.Version) + } + + if want := []string{"v3.1.0", "v3.0.1", "v3.0.0"}; !slices.Equal(got, want) { + t.Fatalf("versions = %v, want %v", got, want) + } + + // v3.0.1 exists in both; the fetched entry (Windows asset) wins. + if releases[1].Assets[0].OS != "windows" { + t.Errorf("v3.0.1 assets = %+v, want the fetched windows asset", releases[1].Assets) + } +} + +func TestListVersionsFetchDisabled(t *testing.T) { + t.Parallel() + + tmp := buildOfflineFixture(t) + + cfg := newTestConfig(tmp) + cfg.DisableFetch = true + + svc := New(cfg, logger.New("error")) + svc.githubBase = newReleasesServer(t, `[]`, false).URL + + releases, err := svc.ListVersions(context.Background()) + if err != nil { + t.Fatal(err) + } + + if len(releases) != 1 || releases[0].Version != "v3.0.1" { + t.Fatalf("releases = %+v, want only local v3.0.1", releases) + } +} + +func TestBuildPackageFetchDisabled(t *testing.T) { + t.Parallel() + + tmp := buildOfflineFixture(t) + + cfg := newTestConfig(tmp) + cfg.DisableFetch = true + + svc := New(cfg, logger.New("error")) + svc.githubBase = newReleasesServer(t, `[]`, false).URL + + _, _, err := svc.BuildPackage(context.Background(), dto.PackageRequest{ + Command: "deactivate", + Version: "v3.0.1", + OS: "linux", + Arch: "x64", + Auth: dto.PackageAuth{Mode: "none"}, + }, "") + if err != nil { + t.Fatal(err) + } +} diff --git a/internal/usecase/packaging/scheme_test.go b/internal/usecase/packaging/scheme_test.go new file mode 100644 index 000000000..ba8688ed2 --- /dev/null +++ b/internal/usecase/packaging/scheme_test.go @@ -0,0 +1,134 @@ +package packaging + +import ( + "strings" + "testing" + + "github.com/device-management-toolkit/console/config" + dto "github.com/device-management-toolkit/console/internal/entity/dto/v1" + "github.com/device-management-toolkit/console/pkg/logger" +) + +// The derived fallback URL must match the scheme the listener actually serves, +// or the generated config points rpc-go at the wrong protocol. +func TestBuildConfigInputsDerivedSchemeFollowsTLS(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + tlsEnabled bool + certFile string + wantPrefix string + wantSkipCerts bool + }{ + {"tls off yields http", false, "", "http://localhost:8181", false}, + {"tls on yields https and skips cert check for self-signed", true, "", "https://localhost:8181", true}, + {"configured cert keeps verification on", true, "/etc/console/tls.crt", "https://localhost:8181", false}, + } + + for _, tt := range tests { + tt := tt + + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + cfg := newTestConfig("") + cfg.TLS = config.TLS{Enabled: tt.tlsEnabled, CertFile: tt.certFile} + + svc := New(cfg, logger.New("error")) + + in, err := svc.buildConfigInputs(dto.PackageRequest{ + Command: "activate", + Auth: dto.PackageAuth{Mode: "userpass"}, + }, "") + if err != nil { + t.Fatal(err) + } + + if !strings.HasPrefix(in.AuthEndpoint, tt.wantPrefix) { + t.Errorf("AuthEndpoint = %q, want prefix %q", in.AuthEndpoint, tt.wantPrefix) + } + + if in.SkipCertCheck != tt.wantSkipCerts { + t.Errorf("SkipCertCheck = %v, want %v", in.SkipCertCheck, tt.wantSkipCerts) + } + }) + } +} + +// A server URL from the request still points at this listener, so a +// self-signed cert must keep the cert check off. +func TestBuildConfigInputsServerURLFollowsTLSCert(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + certFile string + wantSkipCerts bool + }{ + {"self-signed cert skips cert check", "", true}, + {"configured cert keeps verification on", "/etc/console/tls.crt", false}, + } + + for _, tt := range tests { + tt := tt + + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + cfg := newTestConfig("") + cfg.TLS = config.TLS{Enabled: true, CertFile: tt.certFile} + + svc := New(cfg, logger.New("error")) + + in, err := svc.buildConfigInputs(dto.PackageRequest{ + Command: "activate", + Auth: dto.PackageAuth{Mode: "userpass"}, + ServerURL: "https://console.example.com", + }, "") + if err != nil { + t.Fatal(err) + } + + if in.AuthEndpoint != "https://console.example.com/api/v1/authorize" { + t.Errorf("AuthEndpoint = %q", in.AuthEndpoint) + } + + if in.SkipCertCheck != tt.wantSkipCerts { + t.Errorf("SkipCertCheck = %v, want %v", in.SkipCertCheck, tt.wantSkipCerts) + } + }) + } +} + +// The request tenant must reach the generated config, or a package built by one +// tenant resolves its profile against another tenant's data. +func TestRenderConfigCarriesTenant(t *testing.T) { + t.Parallel() + + svc := New(newTestConfig(""), logger.New("error")) + + req := dto.PackageRequest{ + Command: "activate", + Auth: dto.PackageAuth{Mode: "userpass", Username: "u", Password: "p"}, + Profile: "profile1", + } + + in, err := svc.buildConfigInputs(req, "acme-corp") + if err != nil { + t.Fatal(err) + } + + if in.TenantID != "acme-corp" { + t.Fatalf("TenantID = %q, want %q", in.TenantID, "acme-corp") + } + + out, err := renderConfig(req, in) + if err != nil { + t.Fatal(err) + } + + if !strings.Contains(string(out), "tenantid: acme-corp") { + t.Errorf("rendered config missing tenantid:\n%s", out) + } +} diff --git a/internal/usecase/packaging/token.go b/internal/usecase/packaging/token.go new file mode 100644 index 000000000..a15677597 --- /dev/null +++ b/internal/usecase/packaging/token.go @@ -0,0 +1,62 @@ +package packaging + +import ( + "errors" + "fmt" + "time" + + "github.com/golang-jwt/jwt/v5" +) + +const ( + // defaultTokenTTL is how long a minted rpc-go auth token stays valid when the + // request does not ask for a lifetime. + defaultTokenTTL = time.Hour + // defaultMaxTokenTTL caps a requested lifetime when the server sets no maximum. + defaultMaxTokenTTL = 24 * time.Hour +) + +var ( + // ErrTokenTTLInvalid indicates a token lifetime that is not a positive duration. + ErrTokenTTLInvalid = errors.New("invalid token lifetime") + // ErrTokenTTLTooLong indicates a token lifetime above the server maximum. + ErrTokenTTLTooLong = errors.New("token lifetime exceeds the server maximum") + // ErrTokenModeUnsupported indicates token auth under OIDC, where Console-minted tokens are not accepted. + ErrTokenModeUnsupported = errors.New("token auth mode is not supported when OIDC is configured") +) + +// resolveTokenTTL turns a requested lifetime into the duration to mint with. +// An empty request keeps defaultTokenTTL; a zero maxTTL falls back to +// defaultMaxTokenTTL so deployments predating the setting still have a ceiling. +func resolveTokenTTL(requested string, maxTTL time.Duration) (time.Duration, error) { + if maxTTL <= 0 { + maxTTL = defaultMaxTokenTTL + } + + if requested == "" { + return defaultTokenTTL, nil + } + + ttl, err := time.ParseDuration(requested) + if err != nil || ttl <= 0 { + return 0, fmt.Errorf("%w: %q", ErrTokenTTLInvalid, requested) + } + + if ttl > maxTTL { + return 0, fmt.Errorf("%w of %s: %q", ErrTokenTTLTooLong, maxTTL, requested) + } + + return ttl, nil +} + +// mintToken issues an HS256 JWT signed with the given key, mirroring the +// login route's token issuance. rpc-go uses this as its bearer auth-token. +func mintToken(jwtKey string, ttl time.Duration) (string, error) { + claims := jwt.RegisteredClaims{ + ExpiresAt: jwt.NewNumericDate(time.Now().Add(ttl)), + } + + token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims) + + return token.SignedString([]byte(jwtKey)) +} diff --git a/internal/usecase/packaging/token_test.go b/internal/usecase/packaging/token_test.go new file mode 100644 index 000000000..ab65c08b1 --- /dev/null +++ b/internal/usecase/packaging/token_test.go @@ -0,0 +1,106 @@ +package packaging + +import ( + "errors" + "testing" + "time" + + "github.com/golang-jwt/jwt/v5" +) + +func TestMintToken(t *testing.T) { + t.Parallel() + + const key = "test-key" + + tokenString, err := mintToken(key, defaultTokenTTL) + if err != nil { + t.Fatal(err) + } + + if tokenString == "" { + t.Fatal("expected a non-empty token") + } + + claims := &jwt.RegisteredClaims{} + + parsed, err := jwt.ParseWithClaims(tokenString, claims, func(tok *jwt.Token) (interface{}, error) { + if _, ok := tok.Method.(*jwt.SigningMethodHMAC); !ok { + return nil, jwt.ErrSignatureInvalid + } + + return []byte(key), nil + }) + if err != nil { + t.Fatal(err) + } + + if !parsed.Valid { + t.Fatal("expected a valid token") + } + + if claims.ExpiresAt == nil || !claims.ExpiresAt.After(time.Now()) { + t.Fatal("expected an expiry in the future") + } +} + +func TestResolveTokenTTL(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + requested string + maxTTL time.Duration + want time.Duration + wantErr error + }{ + {"empty falls back to the default", "", 24 * time.Hour, defaultTokenTTL, nil}, + {"honors a requested lifetime", "15m", 24 * time.Hour, 15 * time.Minute, nil}, + {"honors a lifetime equal to the cap", "24h", 24 * time.Hour, 24 * time.Hour, nil}, + {"zero cap falls back to the default cap", "24h", 0, 24 * time.Hour, nil}, + {"rejects a lifetime above the cap", "24h", time.Hour, 0, ErrTokenTTLTooLong}, + {"rejects an unparseable lifetime", "fortnight", 24 * time.Hour, 0, ErrTokenTTLInvalid}, + {"rejects a non-positive lifetime", "0s", 24 * time.Hour, 0, ErrTokenTTLInvalid}, + } + + for _, tt := range tests { + tt := tt + + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got, err := resolveTokenTTL(tt.requested, tt.maxTTL) + if !errors.Is(err, tt.wantErr) { + t.Fatalf("error = %v, want %v", err, tt.wantErr) + } + + if got != tt.want { + t.Errorf("ttl = %v, want %v", got, tt.want) + } + }) + } +} + +func TestMintTokenUsesRequestedTTL(t *testing.T) { + t.Parallel() + + const key = "test-key" + + tokenString, err := mintToken(key, 15*time.Minute) + if err != nil { + t.Fatal(err) + } + + claims := &jwt.RegisteredClaims{} + + if _, err := jwt.ParseWithClaims(tokenString, claims, func(_ *jwt.Token) (interface{}, error) { + return []byte(key), nil + }); err != nil { + t.Fatal(err) + } + + got := time.Until(claims.ExpiresAt.Time) + if got < 14*time.Minute || got > 15*time.Minute+time.Minute { + t.Errorf("expiry in %v, want roughly 15m", got) + } +} From 4c348a136c3d41f823b6f6f5ee6b737b93d45e64 Mon Sep 17 00:00:00 2001 From: Mike Date: Fri, 18 Sep 2026 16:27:05 -0700 Subject: [PATCH 3/3] feat(api): add download-rpc package endpoints - GET /api/package/rpc-versions lists the rpc-go releases available to package - POST /api/package returns a zip with the rpc-go binary and a config.yaml pointing at this Console - Map missing assets to 404 and unsafe versions and out-of-range token lifetimes to 400 - Declare both routes in OpenAPI and add Postman requests --- .../console_rps_apis.postman_collection.json | 400 ++++++++++++++++++ internal/controller/httpapi/router.go | 3 + internal/controller/httpapi/v1/error.go | 11 + internal/controller/httpapi/v1/package.go | 80 ++++ .../controller/httpapi/v1/package_test.go | 285 +++++++++++++ internal/controller/openapi/adapter.go | 3 + internal/controller/openapi/package.go | 36 ++ 7 files changed, 818 insertions(+) create mode 100644 internal/controller/httpapi/v1/package.go create mode 100644 internal/controller/httpapi/v1/package_test.go create mode 100644 internal/controller/openapi/package.go diff --git a/integration-test/collections/console_rps_apis.postman_collection.json b/integration-test/collections/console_rps_apis.postman_collection.json index 282b77fbb..68c38e0d4 100644 --- a/integration-test/collections/console_rps_apis.postman_collection.json +++ b/integration-test/collections/console_rps_apis.postman_collection.json @@ -8931,6 +8931,406 @@ "response": [] } ] + }, + { + "name": "Package", + "item": [ + { + "name": "Get RPC Versions", + "event": [ + { + "listen": "test", + "script": { + "exec": [ + "pm.test(\"Status code is 200\", function () {", + " pm.response.to.have.status(200);", + "});", + "", + "pm.test(\"Returns at most five releases\", function () {", + " pm.expect(pm.response.json().length).to.be.at.most(5);", + "});", + "", + "pm.test(\"Every release is v3 or above and lists assets\", function () {", + " pm.response.json().forEach(function (release) {", + " pm.expect(release.version).to.match(/^v?[3-9]\\d*\\./);", + " pm.expect(release.assets).to.be.an(\"array\");", + " });", + "});", + "", + "var releases = pm.response.json();", + "if (releases.length > 0 && releases[0].assets.length > 0) {", + " pm.collectionVariables.set(\"rpcVersion\", releases[0].version);", + " pm.collectionVariables.set(\"rpcOS\", releases[0].assets[0].os);", + " pm.collectionVariables.set(\"rpcArch\", releases[0].assets[0].arch);", + "}", + "" + ], + "type": "text/javascript" + } + } + ], + "request": { + "method": "GET", + "header": [ + { + "key": "Content-Type", + "value": "application/json" + } + ], + "url": { + "raw": "{{protocol}}://{{host}}/api/package/rpc-versions", + "protocol": "{{protocol}}", + "host": [ + "{{host}}" + ], + "path": [ + "api", + "package", + "rpc-versions" + ] + } + }, + "response": [] + }, + { + "name": "Build Package (Activate)", + "event": [ + { + "listen": "test", + "script": { + "exec": [ + "pm.test(\"Status code is 200\", function () {", + " pm.response.to.have.status(200);", + "});", + "", + "pm.test(\"Responds with a zip attachment\", function () {", + " pm.expect(pm.response.headers.get(\"Content-Type\")).to.include(\"application/zip\");", + " pm.expect(pm.response.headers.get(\"Content-Disposition\")).to.include(\".zip\");", + "});", + "" + ], + "type": "text/javascript" + } + } + ], + "request": { + "method": "POST", + "header": [ + { + "key": "Content-Type", + "value": "application/json" + } + ], + "body": { + "mode": "raw", + "raw": "{\n \"command\": \"activate\",\n \"version\": \"{{rpcVersion}}\",\n \"os\": \"{{rpcOS}}\",\n \"arch\": \"{{rpcArch}}\",\n \"auth\": {\n \"mode\": \"userpass\",\n \"username\": \"standalone\",\n \"password\": \"G@ppm0ym\"\n },\n \"profile\": \"profile1\",\n \"tokenTtl\": \"1h\"\n}", + "options": { + "raw": { + "language": "json" + } + } + }, + "url": { + "raw": "{{protocol}}://{{host}}/api/package", + "protocol": "{{protocol}}", + "host": [ + "{{host}}" + ], + "path": [ + "api", + "package" + ] + } + }, + "response": [] + }, + { + "name": "Build Package (Tenant Scoped)", + "event": [ + { + "listen": "test", + "script": { + "exec": [ + "pm.test(\"Status code is 200\", function () {", + " pm.response.to.have.status(200);", + "});", + "" + ], + "type": "text/javascript" + } + } + ], + "request": { + "method": "POST", + "header": [ + { + "key": "Content-Type", + "value": "application/json" + }, + { + "key": "x-tenant-id", + "value": "{{tenantA}}" + } + ], + "body": { + "mode": "raw", + "raw": "{\n \"command\": \"activate\",\n \"version\": \"{{rpcVersion}}\",\n \"os\": \"{{rpcOS}}\",\n \"arch\": \"{{rpcArch}}\",\n \"auth\": {\n \"mode\": \"userpass\",\n \"username\": \"standalone\",\n \"password\": \"G@ppm0ym\"\n },\n \"profile\": \"profile1\"\n}", + "options": { + "raw": { + "language": "json" + } + } + }, + "url": { + "raw": "{{protocol}}://{{host}}/api/package", + "protocol": "{{protocol}}", + "host": [ + "{{host}}" + ], + "path": [ + "api", + "package" + ] + } + }, + "response": [] + }, + { + "name": "Build Package (Deactivate)", + "event": [ + { + "listen": "test", + "script": { + "exec": [ + "pm.test(\"Status code is 200\", function () {", + " pm.response.to.have.status(200);", + "});", + "" + ], + "type": "text/javascript" + } + } + ], + "request": { + "method": "POST", + "header": [ + { + "key": "Content-Type", + "value": "application/json" + } + ], + "body": { + "mode": "raw", + "raw": "{\n \"command\": \"deactivate\",\n \"version\": \"{{rpcVersion}}\",\n \"os\": \"{{rpcOS}}\",\n \"arch\": \"{{rpcArch}}\",\n \"auth\": {\n \"mode\": \"userpass\",\n \"username\": \"standalone\",\n \"password\": \"G@ppm0ym\"\n }\n}", + "options": { + "raw": { + "language": "json" + } + } + }, + "url": { + "raw": "{{protocol}}://{{host}}/api/package", + "protocol": "{{protocol}}", + "host": [ + "{{host}}" + ], + "path": [ + "api", + "package" + ] + } + }, + "response": [] + }, + { + "name": "Build Package (No Credentials)", + "event": [ + { + "listen": "test", + "script": { + "exec": [ + "pm.test(\"Status code is 200\", function () {", + " pm.response.to.have.status(200);", + "});", + "" + ], + "type": "text/javascript" + } + } + ], + "request": { + "method": "POST", + "header": [ + { + "key": "Content-Type", + "value": "application/json" + } + ], + "body": { + "mode": "raw", + "raw": "{\n \"command\": \"activate\",\n \"version\": \"{{rpcVersion}}\",\n \"os\": \"{{rpcOS}}\",\n \"arch\": \"{{rpcArch}}\",\n \"auth\": {\n \"mode\": \"none\"\n },\n \"profile\": \"profile1\"\n}", + "options": { + "raw": { + "language": "json" + } + } + }, + "url": { + "raw": "{{protocol}}://{{host}}/api/package", + "protocol": "{{protocol}}", + "host": [ + "{{host}}" + ], + "path": [ + "api", + "package" + ] + } + }, + "response": [] + }, + { + "name": "Build Package (Activate Without Profile - 400)", + "event": [ + { + "listen": "test", + "script": { + "exec": [ + "pm.test(\"Status code is 400\", function () {", + " pm.response.to.have.status(400);", + "});", + "" + ], + "type": "text/javascript" + } + } + ], + "request": { + "method": "POST", + "header": [ + { + "key": "Content-Type", + "value": "application/json" + } + ], + "body": { + "mode": "raw", + "raw": "{\n \"command\": \"activate\",\n \"version\": \"{{rpcVersion}}\",\n \"os\": \"{{rpcOS}}\",\n \"arch\": \"{{rpcArch}}\",\n \"auth\": {\n \"mode\": \"userpass\",\n \"username\": \"standalone\",\n \"password\": \"G@ppm0ym\"\n }\n}", + "options": { + "raw": { + "language": "json" + } + } + }, + "url": { + "raw": "{{protocol}}://{{host}}/api/package", + "protocol": "{{protocol}}", + "host": [ + "{{host}}" + ], + "path": [ + "api", + "package" + ] + } + }, + "response": [] + }, + { + "name": "Build Package (Unknown Version - 404)", + "event": [ + { + "listen": "test", + "script": { + "exec": [ + "pm.test(\"Status code is 404\", function () {", + " pm.response.to.have.status(404);", + "});", + "" + ], + "type": "text/javascript" + } + } + ], + "request": { + "method": "POST", + "header": [ + { + "key": "Content-Type", + "value": "application/json" + } + ], + "body": { + "mode": "raw", + "raw": "{\n \"command\": \"activate\",\n \"version\": \"v99.99.99\",\n \"os\": \"linux\",\n \"arch\": \"x86_64\",\n \"auth\": {\n \"mode\": \"userpass\",\n \"username\": \"standalone\",\n \"password\": \"G@ppm0ym\"\n },\n \"profile\": \"profile1\"\n}", + "options": { + "raw": { + "language": "json" + } + } + }, + "url": { + "raw": "{{protocol}}://{{host}}/api/package", + "protocol": "{{protocol}}", + "host": [ + "{{host}}" + ], + "path": [ + "api", + "package" + ] + } + }, + "response": [] + }, + { + "name": "Build Package (Path Traversal Version - 400)", + "event": [ + { + "listen": "test", + "script": { + "exec": [ + "pm.test(\"Status code is 400\", function () {", + " pm.response.to.have.status(400);", + "});", + "" + ], + "type": "text/javascript" + } + } + ], + "request": { + "method": "POST", + "header": [ + { + "key": "Content-Type", + "value": "application/json" + } + ], + "body": { + "mode": "raw", + "raw": "{\n \"command\": \"activate\",\n \"version\": \"../../etc\",\n \"os\": \"linux\",\n \"arch\": \"x86_64\",\n \"auth\": {\n \"mode\": \"userpass\",\n \"username\": \"standalone\",\n \"password\": \"G@ppm0ym\"\n },\n \"profile\": \"profile1\"\n}", + "options": { + "raw": { + "language": "json" + } + } + }, + "url": { + "raw": "{{protocol}}://{{host}}/api/package", + "protocol": "{{protocol}}", + "host": [ + "{{host}}" + ], + "path": [ + "api", + "package" + ] + } + }, + "response": [] + } + ] } ], "auth": { diff --git a/internal/controller/httpapi/router.go b/internal/controller/httpapi/router.go index 9d98cf0af..6c60dee24 100644 --- a/internal/controller/httpapi/router.go +++ b/internal/controller/httpapi/router.go @@ -17,6 +17,7 @@ import ( openapi "github.com/device-management-toolkit/console/internal/controller/openapi" dto "github.com/device-management-toolkit/console/internal/entity/dto/v1" "github.com/device-management-toolkit/console/internal/usecase" + "github.com/device-management-toolkit/console/internal/usecase/packaging" "github.com/device-management-toolkit/console/pkg/logger" ) @@ -91,6 +92,8 @@ func NewRouter(handler *gin.Engine, l logger.Interface, t usecase.Usecases, cfg { v2.NewAmtRoutes(h3, t.Devices, l) } + + v1.NewPackageRoutes(protected, packaging.New(cfg, l), l) } func registerCustomValidators(l logger.Interface) { diff --git a/internal/controller/httpapi/v1/error.go b/internal/controller/httpapi/v1/error.go index a6152a5ba..88523df1c 100644 --- a/internal/controller/httpapi/v1/error.go +++ b/internal/controller/httpapi/v1/error.go @@ -14,6 +14,7 @@ import ( "github.com/device-management-toolkit/console/internal/usecase/devices" wsmanAPI "github.com/device-management-toolkit/console/internal/usecase/devices/wsman" "github.com/device-management-toolkit/console/internal/usecase/domains" + "github.com/device-management-toolkit/console/internal/usecase/packaging" "github.com/device-management-toolkit/console/internal/usecase/profiles" "github.com/device-management-toolkit/console/internal/usecase/sqldb" ) @@ -141,6 +142,16 @@ func handleSentinelErrors(c *gin.Context, err error) bool { msg := wsmanAPI.ErrCIRADeviceNotConnected.Error() c.AbortWithStatusJSON(http.StatusServiceUnavailable, response{Error: msg, Message: msg}) + return true + case errors.Is(err, packaging.ErrAssetNotFound): + msg := err.Error() + c.AbortWithStatusJSON(http.StatusNotFound, response{Error: msg, Message: msg}) + + return true + case errors.Is(err, packaging.ErrUnsafeVersion): + msg := err.Error() + c.AbortWithStatusJSON(http.StatusBadRequest, response{Error: msg, Message: msg}) + return true } diff --git a/internal/controller/httpapi/v1/package.go b/internal/controller/httpapi/v1/package.go new file mode 100644 index 000000000..a30d8fa0c --- /dev/null +++ b/internal/controller/httpapi/v1/package.go @@ -0,0 +1,80 @@ +package v1 + +import ( + "errors" + "net/http" + + "github.com/gin-gonic/gin" + + "github.com/device-management-toolkit/console/internal/controller/httpapi/middleware" + dto "github.com/device-management-toolkit/console/internal/entity/dto/v1" + "github.com/device-management-toolkit/console/internal/usecase/packaging" + "github.com/device-management-toolkit/console/pkg/consoleerrors" + "github.com/device-management-toolkit/console/pkg/logger" +) + +const unknownContentLength int64 = -1 + +var errValidationPackage = dto.NotValidError{Console: consoleerrors.CreateConsoleError("PackageAPI")} + +type packageRoutes struct { + t packaging.Feature + l logger.Interface +} + +// NewPackageRoutes registers the Download RPC endpoints under the given group. +func NewPackageRoutes(handler *gin.RouterGroup, t packaging.Feature, l logger.Interface) { + r := &packageRoutes{t: t, l: l} + + h := handler.Group("/package") + { + h.GET("/rpc-versions", r.versions) + h.POST("", r.build) + } +} + +func (r *packageRoutes) versions(c *gin.Context) { + releases, err := r.t.ListVersions(c.Request.Context()) + if err != nil { + r.l.Error(err, "http - v1 - package - rpc-versions") + ErrorResponse(c, err) + + return + } + + c.JSON(http.StatusOK, releases) +} + +func (r *packageRoutes) build(c *gin.Context) { + var req dto.PackageRequest + if err := c.ShouldBindJSON(&req); err != nil { + validationErr := errValidationPackage.Wrap("build", "ShouldBindJSON", err) + ErrorResponse(c, validationErr) + + return + } + + reader, filename, err := r.t.BuildPackage(c.Request.Context(), req, middleware.TenantID(c)) + if err != nil { + if errors.Is(err, packaging.ErrTokenTTLTooLong) || errors.Is(err, packaging.ErrTokenTTLInvalid) { + ErrorResponse(c, errValidationPackage.Wrap("build", "tokenTtl", err)) + + return + } + + if errors.Is(err, packaging.ErrTokenModeUnsupported) { + ErrorResponse(c, errValidationPackage.Wrap("build", "auth.mode", err)) + + return + } + + r.l.Error(err, "http - v1 - package - build") + ErrorResponse(c, err) + + return + } + + c.DataFromReader(http.StatusOK, unknownContentLength, "application/zip", reader, map[string]string{ + "Content-Disposition": `attachment; filename="` + filename + `"`, + }) +} diff --git a/internal/controller/httpapi/v1/package_test.go b/internal/controller/httpapi/v1/package_test.go new file mode 100644 index 000000000..47bb03ec2 --- /dev/null +++ b/internal/controller/httpapi/v1/package_test.go @@ -0,0 +1,285 @@ +package v1 + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + + "github.com/device-management-toolkit/console/internal/controller/httpapi/middleware" + dto "github.com/device-management-toolkit/console/internal/entity/dto/v1" + "github.com/device-management-toolkit/console/internal/usecase/packaging" + "github.com/device-management-toolkit/console/pkg/logger" +) + +type stubPackaging struct { + releases []dto.RPCRelease + zip []byte + err error + tenantID string +} + +func (s *stubPackaging) ListVersions(_ context.Context) ([]dto.RPCRelease, error) { + return s.releases, s.err +} + +func (s *stubPackaging) BuildPackage(_ context.Context, _ dto.PackageRequest, tenantID string) (io.Reader, string, error) { + s.tenantID = tenantID + + if s.err != nil { + return nil, "", s.err + } + + return bytes.NewReader(s.zip), "rpc-activate-linux-x86_64.zip", nil +} + +func newPackageEngine(stub *stubPackaging) *gin.Engine { + log := logger.New("error") + engine := gin.New() + + // Mirror router.go, which resolves the tenant on the protected group before + // the package routes are mounted. + group := engine.Group("/api") + group.Use(middleware.ResolveTenant(log)) + + NewPackageRoutes(group, stub, log) + + return engine +} + +func TestPackageRoutes(t *testing.T) { + t.Parallel() + + t.Run("GET rpc-versions returns 200 with releases", func(t *testing.T) { + t.Parallel() + + releases := []dto.RPCRelease{ + {Version: "v1.2.3", Assets: []dto.RPCAsset{{OS: "linux", Arch: "x86_64"}}}, + } + engine := newPackageEngine(&stubPackaging{releases: releases}) + + req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, "/api/package/rpc-versions", http.NoBody) + require.NoError(t, err) + + w := httptest.NewRecorder() + engine.ServeHTTP(w, req) + + require.Equal(t, http.StatusOK, w.Code) + + wantJSON, err := json.Marshal(releases) + require.NoError(t, err) + require.Equal(t, string(wantJSON), w.Body.String()) + }) + + t.Run("POST package with invalid body returns 400", func(t *testing.T) { + t.Parallel() + + // Malformed JSON triggers a JSON-decode error from ShouldBindJSON, + // which the handler wraps as a NotValidError → 400 Bad Request. + // (gin.DisableBindValidation is set in init() so struct-tag validation + // is not active in tests; a JSON-syntax error is the reliable way to + // exercise the 400 path.) + body := `{not valid json` + engine := newPackageEngine(&stubPackaging{}) + + req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, "/api/package", bytes.NewBufferString(body)) + require.NoError(t, err) + req.Header.Set("Content-Type", "application/json") + + w := httptest.NewRecorder() + engine.ServeHTTP(w, req) + + require.Equal(t, http.StatusBadRequest, w.Code) + }) + + t.Run("POST package with valid body returns 200 zip", func(t *testing.T) { + t.Parallel() + + zipData := []byte("PK\x03\x04fake-zip-content") + engine := newPackageEngine(&stubPackaging{zip: zipData}) + + reqBody := dto.PackageRequest{ + Command: "activate", + Version: "v1.2.3", + OS: "linux", + Arch: "x86_64", + Auth: dto.PackageAuth{Mode: "token"}, + } + + bodyBytes, err := json.Marshal(reqBody) + require.NoError(t, err) + + req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, "/api/package", bytes.NewBuffer(bodyBytes)) + require.NoError(t, err) + req.Header.Set("Content-Type", "application/json") + + w := httptest.NewRecorder() + engine.ServeHTTP(w, req) + + require.Equal(t, http.StatusOK, w.Code) + require.Contains(t, w.Header().Get("Content-Type"), "application/zip") + require.Equal(t, zipData, w.Body.Bytes()) + }) + + t.Run("POST package returns 400 when the token lifetime exceeds the server maximum", func(t *testing.T) { + t.Parallel() + + engine := newPackageEngine(&stubPackaging{err: packaging.ErrTokenTTLTooLong}) + + reqBody := dto.PackageRequest{ + Command: "activate", + Version: "v1.2.3", + OS: "linux", + Arch: "x86_64", + Auth: dto.PackageAuth{Mode: "token"}, + TokenTTL: "24h", + } + + bodyBytes, err := json.Marshal(reqBody) + require.NoError(t, err) + + req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, "/api/package", bytes.NewBuffer(bodyBytes)) + require.NoError(t, err) + req.Header.Set("Content-Type", "application/json") + + w := httptest.NewRecorder() + engine.ServeHTTP(w, req) + + require.Equal(t, http.StatusBadRequest, w.Code) + }) + + t.Run("GET rpc-versions returns 5xx when ListVersions errors", func(t *testing.T) { + t.Parallel() + + stubErr := errors.New("upstream unavailable") + engine := newPackageEngine(&stubPackaging{err: stubErr}) + + req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, "/api/package/rpc-versions", http.NoBody) + require.NoError(t, err) + + w := httptest.NewRecorder() + engine.ServeHTTP(w, req) + + require.GreaterOrEqual(t, w.Code, http.StatusInternalServerError) + }) + + t.Run("POST package returns 5xx when BuildPackage errors", func(t *testing.T) { + t.Parallel() + + stubErr := errors.New("build failure") + engine := newPackageEngine(&stubPackaging{err: stubErr}) + + reqBody := dto.PackageRequest{ + Command: "activate", + Version: "v1.2.3", + OS: "linux", + Arch: "x86_64", + Auth: dto.PackageAuth{Mode: "token"}, + } + + bodyBytes, err := json.Marshal(reqBody) + require.NoError(t, err) + + req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, "/api/package", bytes.NewBuffer(bodyBytes)) + require.NoError(t, err) + req.Header.Set("Content-Type", "application/json") + + w := httptest.NewRecorder() + engine.ServeHTTP(w, req) + + require.GreaterOrEqual(t, w.Code, http.StatusInternalServerError) + }) +} + +// postPackage issues a build request against a stub and returns the recorder. +func postPackage(t *testing.T, stub *stubPackaging, headers map[string]string) *httptest.ResponseRecorder { + t.Helper() + + engine := newPackageEngine(stub) + + req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, "/api/package", bytes.NewBufferString(validPackageBody)) + require.NoError(t, err) + req.Header.Set("Content-Type", "application/json") + + for k, v := range headers { + req.Header.Set(k, v) + } + + w := httptest.NewRecorder() + engine.ServeHTTP(w, req) + + return w +} + +const validPackageBody = `{"command":"activate","version":"v3.0.1","os":"linux","arch":"x86_64",` + + `"auth":{"mode":"userpass","username":"u","password":"p"},"profile":"p1"}` + +// The resolved tenant must reach the use case; without it a package built by one +// tenant carries no tenant scope and resolves against the default tenant's data. +func TestPackageBuildPropagatesTenant(t *testing.T) { + t.Parallel() + + stub := &stubPackaging{zip: []byte("zip-bytes")} + + w := postPackage(t, stub, map[string]string{"x-tenant-id": "acme-corp"}) + + require.Equal(t, http.StatusOK, w.Code) + require.Equal(t, "acme-corp", stub.tenantID) +} + +func TestPackageBuildDefaultTenantWhenHeaderAbsent(t *testing.T) { + t.Parallel() + + stub := &stubPackaging{zip: []byte("zip-bytes")} + + w := postPackage(t, stub, nil) + + require.Equal(t, http.StatusOK, w.Code) + require.Empty(t, stub.tenantID) +} + +// Bad input must not read as a server fault. +func TestPackageBuildMapsUserErrors(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + err error + wantCode int + }{ + {"unknown asset is not found", packaging.ErrAssetNotFound, http.StatusNotFound}, + {"traversal version is a bad request", packaging.ErrUnsafeVersion, http.StatusBadRequest}, + {"token mode under OIDC is a bad request", packaging.ErrTokenModeUnsupported, http.StatusBadRequest}, + } + + for _, tt := range tests { + tt := tt + + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + w := postPackage(t, &stubPackaging{err: tt.err}, nil) + + require.Equal(t, tt.wantCode, w.Code) + }) + } +} + +// Wrapped errors must map the same way — BuildPackage wraps with %w. +func TestPackageBuildMapsWrappedAssetNotFound(t *testing.T) { + t.Parallel() + + wrapped := fmt.Errorf("resolve asset: %w", packaging.ErrAssetNotFound) + + w := postPackage(t, &stubPackaging{err: wrapped}, nil) + + require.Equal(t, http.StatusNotFound, w.Code) +} diff --git a/internal/controller/openapi/adapter.go b/internal/controller/openapi/adapter.go index daa0b5825..949800946 100644 --- a/internal/controller/openapi/adapter.go +++ b/internal/controller/openapi/adapter.go @@ -102,6 +102,9 @@ func (f *FuegoAdapter) RegisterRoutes() { // Server features f.RegisterServerRoutes() + + // Download RPC packaging + f.RegisterPackageRoutes() } // Generates OpenAPI specification as JSON. diff --git a/internal/controller/openapi/package.go b/internal/controller/openapi/package.go new file mode 100644 index 000000000..28dfec59a --- /dev/null +++ b/internal/controller/openapi/package.go @@ -0,0 +1,36 @@ +package openapi + +import ( + "net/http" + + "github.com/go-fuego/fuego" + + dto "github.com/device-management-toolkit/console/internal/entity/dto/v1" +) + +// RegisterPackageRoutes declares the Download RPC endpoints. They sit under +// /api/package rather than a versioned group, matching the Gin routes. +func (f *FuegoAdapter) RegisterPackageRoutes() { + fuego.Get(f.server, "/api/package/rpc-versions", f.listRPCVersions, + fuego.OptionTags("Package"), + fuego.OptionSummary("List RPC Versions"), + fuego.OptionDescription("List the most recent rpc-go releases (v3 and above) available for packaging, with the OS/arch builds each one publishes"), + protectedRouteOptions(), + ) + + fuego.Post(f.server, "/api/package", f.buildPackage, + fuego.OptionTags("Package"), + fuego.OptionSummary("Build RPC Package"), + fuego.OptionDescription("Build a zip containing the requested rpc-go binary and a generated config.yaml pointing at this Console. The response is the zip itself, not JSON."), + fuego.OptionAddResponse(http.StatusOK, "OK", fuego.Response{Type: "", ContentTypes: []string{"application/zip"}}), + protectedRouteOptions(), + ) +} + +func (f *FuegoAdapter) listRPCVersions(_ fuego.ContextNoBody) ([]dto.RPCRelease, error) { + return nil, nil +} + +func (f *FuegoAdapter) buildPackage(_ fuego.ContextWithBody[dto.PackageRequest]) (string, error) { + return "", nil +}