diff --git a/e2e-docker/virtual_test.go b/e2e-docker/virtual_test.go
index 42796cc..bc3309b 100644
--- a/e2e-docker/virtual_test.go
+++ b/e2e-docker/virtual_test.go
@@ -3,9 +3,17 @@
package e2edocker
import (
+ "bytes"
+ "compress/gzip"
+ "encoding/xml"
+ "fmt"
+ "io"
"net/http"
+ "os"
+ "os/exec"
"strings"
"testing"
+ "time"
)
// TestVirtualPyPIMerge uploads different packages to two pypi locals and
@@ -52,3 +60,79 @@ func TestVirtualHelmMerge(t *testing.T) {
t.Fatalf("merged helm index missing a member chart (want alpha and beta): %s", s)
}
}
+
+// TestVirtualRPMMerge merges a remote rpm repo and a local rpm repo carrying the
+// same package: the first member wins the duplicate, its package download
+// routes through the virtual, and a real dnf can consume the merged repo.
+func TestVirtualRPMMerge(t *testing.T) {
+ createRepo(t, `{"name":"vrpm-remote","package_type":"rpm","repo_type":"remote","base_url":"`+mockUpstream()+`/rpm-mirror","stale_on_error":true}`)
+ createRepo(t, `{"name":"vrpm-local","package_type":"rpm","repo_type":"local"}`)
+ defer deleteRepo(t, "vrpm-remote")
+ defer deleteRepo(t, "vrpm-local")
+
+ pkg := fixtureBytes(t, "rpmrepo/Packages/e2e-testpkg-1.0-1.noarch.rpm")
+ uploadFile(t, "vrpm-local", "e2e-testpkg-1.0-1.noarch.rpm", pkg, "application/x-rpm")
+
+ createVirtual(t, `{"name":"vrpm","package_type":"rpm","members":["vrpm-remote","vrpm-local"]}`)
+ defer deleteVirtual(t, "vrpm")
+
+ resp, body := getEventually(t, api("/api/v1/virtual/vrpm/repodata/repomd.xml"), 15*time.Second)
+ if resp.StatusCode != http.StatusOK {
+ t.Fatalf("virtual repomd.xml: status %d: %s", resp.StatusCode, body)
+ }
+ var md struct {
+ Data []struct {
+ Type string `xml:"type,attr"`
+ Location struct {
+ Href string `xml:"href,attr"`
+ } `xml:"location"`
+ } `xml:"data"`
+ }
+ if err := xml.Unmarshal(body, &md); err != nil {
+ t.Fatalf("parse repomd.xml: %v\n%s", err, body)
+ }
+ var primary []byte
+ for _, d := range md.Data {
+ if d.Type == "primary" {
+ resp, gz := doRequest(t, http.MethodGet, api("/api/v1/virtual/vrpm/"+d.Location.Href), nil, "")
+ if resp.StatusCode != http.StatusOK {
+ t.Fatalf("primary: status %d", resp.StatusCode)
+ }
+ r, err := gzip.NewReader(bytes.NewReader(gz))
+ if err != nil {
+ t.Fatal(err)
+ }
+ primary, _ = io.ReadAll(r)
+ }
+ }
+ href := "vrpm-remote/Packages/e2e-testpkg-1.0-1.noarch.rpm"
+ if strings.Count(string(primary), "e2e-testpkg") != 1 || !strings.Contains(string(primary), `href="`+href+`"`) {
+ t.Fatalf("merged primary should hold one e2e-testpkg owned by the first member:\n%s", primary)
+ }
+
+ resp, got := doRequest(t, http.MethodGet, api("/api/v1/virtual/vrpm/"+href), nil, "")
+ if resp.StatusCode != http.StatusOK || !bytes.Equal(got, pkg) {
+ t.Fatalf("package via virtual: status %d, %d bytes (want %d)", resp.StatusCode, len(got), len(pkg))
+ }
+
+ network := os.Getenv("COMPOSE_NETWORK")
+ internal := os.Getenv("ARTIFACTAPI_INTERNAL")
+ if network == "" || internal == "" {
+ t.Skip("COMPOSE_NETWORK/ARTIFACTAPI_INTERNAL not set; skipping real dnf")
+ }
+ if _, err := exec.LookPath("docker"); err != nil {
+ t.Skip("docker not available on the test host")
+ }
+ repoConf := fmt.Sprintf("[vrpm]\nname=vrpm\nbaseurl=%s/api/v1/virtual/vrpm/\nenabled=1\ngpgcheck=0\nmetadata_expire=0\n", strings.TrimRight(internal, "/"))
+ script := "set -euo pipefail; " +
+ "printf '%s' \"$REPO\" > /etc/yum.repos.d/vrpm.repo; " +
+ "dnf -y --disablerepo='*' --enablerepo=vrpm makecache; " +
+ "dnf -y --disablerepo='*' --enablerepo=vrpm list --available; " +
+ "dnf -y --disablerepo='*' --enablerepo=vrpm install e2e-testpkg; " +
+ "rpm -q e2e-testpkg"
+ out, err := exec.Command("docker", "run", "--rm", "--network", network, "-e", "REPO="+repoConf,
+ "rockylinux:9", "bash", "-c", script).CombinedOutput()
+ if err != nil {
+ t.Fatalf("real dnf against rpm virtual failed: %v\n%s", err, out)
+ }
+}
diff --git a/go.mod b/go.mod
index 6333d15..6583252 100644
--- a/go.mod
+++ b/go.mod
@@ -18,6 +18,7 @@ require (
github.com/testcontainers/testcontainers-go/modules/redis v0.42.0
github.com/ulikunitz/xz v0.5.16
golang.org/x/crypto v0.54.0
+ golang.org/x/sync v0.22.0
golang.org/x/time v0.15.0
gopkg.in/yaml.v3 v3.0.1
)
@@ -101,7 +102,6 @@ require (
go.uber.org/atomic v1.11.0 // indirect
go.yaml.in/yaml/v3 v3.0.4 // indirect
golang.org/x/net v0.56.0 // indirect
- golang.org/x/sync v0.22.0 // indirect
golang.org/x/sys v0.47.0 // indirect
golang.org/x/text v0.40.0 // indirect
gopkg.in/ini.v1 v1.67.2 // indirect
diff --git a/internal/api/v1/proxy.go b/internal/api/v1/proxy.go
index bb0474b..e77f43e 100644
--- a/internal/api/v1/proxy.go
+++ b/internal/api/v1/proxy.go
@@ -162,7 +162,21 @@ func (h *ProxyHandler) handleVirtual(w http.ResponseWriter, r *http.Request) {
proxyBaseURL := fmt.Sprintf("%s://%s", scheme(r), r.Host)
+ loc, ok, err := h.virtualEngine.MemberRedirect(r.Context(), *virt, path, r.URL.RawQuery, proxyBaseURL)
+ if err != nil {
+ http.Error(w, err.Error(), http.StatusBadRequest)
+ return
+ }
+ if ok {
+ http.Redirect(w, r, loc, http.StatusFound)
+ return
+ }
+
body, contentType, err := h.virtualEngine.Fetch(r.Context(), *virt, path, proxyBaseURL)
+ if errors.Is(err, virtual.ErrNotFound) {
+ http.Error(w, "not found", http.StatusNotFound)
+ return
+ }
if err != nil {
slog.Error("virtual fetch failed", "virtual", virtualName, "path", path, "error", err)
http.Error(w, "bad gateway", http.StatusBadGateway)
diff --git a/internal/virtual/engine.go b/internal/virtual/engine.go
index bcc352e..d4c52d0 100644
--- a/internal/virtual/engine.go
+++ b/internal/virtual/engine.go
@@ -2,27 +2,58 @@ package virtual
import (
"context"
+ "errors"
"fmt"
"io"
"log/slog"
+ "net/http"
+ "net/http/httptest"
+ "net/url"
+ "slices"
+ "strings"
"sync"
+ "time"
"git.unkin.net/unkin/artifactapi/internal/database"
"git.unkin.net/unkin/artifactapi/internal/provider"
"git.unkin.net/unkin/artifactapi/internal/proxy"
"git.unkin.net/unkin/artifactapi/pkg/models"
+ "golang.org/x/sync/singleflight"
)
type Engine struct {
db *database.DB
proxyEngine *proxy.Engine
+ getRemote func(context.Context, string) (*models.Remote, error)
+ rpmMember func(context.Context, string) (*RPMMember, error)
+
+ mergeTTL time.Duration
+
+ sf singleflight.Group
+ mu sync.Mutex
+ rpm map[string]*rpmGen
}
+// rpmGen holds a virtual's current merge plus the previous one, so a client
+// holding the previous repomd can still fetch its content-hashed data files.
+type rpmGen struct {
+ cur, prev *RPMRepo
+ at time.Time
+}
+
+const rpmMergeTTL = 60 * time.Second
+
func NewEngine(db *database.DB, proxyEngine *proxy.Engine) *Engine {
- return &Engine{db: db, proxyEngine: proxyEngine}
+ e := &Engine{db: db, proxyEngine: proxyEngine, getRemote: db.GetRemote, mergeTTL: rpmMergeTTL}
+ e.rpmMember = e.fetchRPMMember
+ return e
}
func (e *Engine) Fetch(ctx context.Context, virt models.Virtual, path string, proxyBaseURL string) ([]byte, string, error) {
+ if virt.PackageType == models.PackageRPM {
+ return e.fetchRPM(ctx, virt, path)
+ }
+
merger, err := GetMerger(virt.PackageType)
if err != nil {
return nil, "", fmt.Errorf("unsupported virtual type %q: %w", virt.PackageType, err)
@@ -133,3 +164,233 @@ func (e *Engine) fetchLocalIndex(ctx context.Context, remote models.Remote, path
return indexer.GenerateLocalIndex(ctx, e.db, remote.Name, path)
}
+
+var (
+ ErrNotFound = errors.New("not found")
+ ErrBadPath = errors.New("invalid path")
+)
+
+// MemberRedirect maps an rpm virtual package path (/, as written
+// by MergeRPM) to the owning member's route. ok is false for any other path;
+// ErrBadPath means the path is absolute or escapes the member.
+func (e *Engine) MemberRedirect(ctx context.Context, virt models.Virtual, path, rawQuery, proxyBaseURL string) (string, bool, error) {
+ if virt.PackageType != models.PackageRPM || strings.HasPrefix(path, "repodata/") {
+ return "", false, nil
+ }
+ name, rest, found := strings.Cut(path, "/")
+ if !found || !slices.Contains(virt.Members, name) {
+ return "", false, nil
+ }
+ segs, err := memberSegments(rest)
+ if err != nil {
+ return "", false, err
+ }
+ remote, err := e.getRemote(ctx, name)
+ if err != nil {
+ return "", false, nil
+ }
+ loc := fmt.Sprintf("%s/api/v1/%s/%s/%s", strings.TrimRight(proxyBaseURL, "/"), remote.RepoType, url.PathEscape(name), strings.Join(segs, "/"))
+ if rawQuery != "" {
+ loc += "?" + rawQuery
+ }
+ return loc, true, nil
+}
+
+// memberSegments decodes rest and returns its path-escaped segments, rejecting
+// absolute paths and any "." or ".." segment.
+func memberSegments(rest string) ([]string, error) {
+ dec, err := url.PathUnescape(rest)
+ if err != nil || dec == "" || strings.HasPrefix(dec, "/") {
+ return nil, ErrBadPath
+ }
+ segs := strings.Split(dec, "/")
+ for i, s := range segs {
+ if s == "." || s == ".." || strings.Contains(s, "\\") {
+ return nil, ErrBadPath
+ }
+ segs[i] = url.PathEscape(s)
+ }
+ return segs, nil
+}
+
+func (e *Engine) fetchRPM(ctx context.Context, virt models.Virtual, path string) ([]byte, string, error) {
+ if !strings.HasPrefix(path, "repodata/") {
+ return nil, "", ErrNotFound
+ }
+ if path != "repodata/repomd.xml" {
+ if body, ok := e.cachedRPMFile(virt.Name, path); ok {
+ return body, "application/gzip", nil
+ }
+ }
+
+ repo, err := e.mergedRPM(ctx, virt)
+ if err != nil {
+ return nil, "", err
+ }
+ if path == "repodata/repomd.xml" {
+ return repo.Repomd, "application/xml", nil
+ }
+ if body, ok := repo.Files[path]; ok {
+ return body, "application/gzip", nil
+ }
+ return nil, "", ErrNotFound
+}
+
+// mergedRPM returns the virtual's merge, reusing it for mergeTTL and
+// collapsing concurrent merges of the same virtual into one.
+// ponytail: per-replica cache; behind a non-sticky LB a data request landing on
+// another replica re-merges and 404s if a member changed in between. Move to the
+// shared redis cache if that shows up.
+func (e *Engine) mergedRPM(ctx context.Context, virt models.Virtual) (*RPMRepo, error) {
+ e.mu.Lock()
+ g := e.rpm[virt.Name]
+ if g != nil && time.Since(g.at) < e.mergeTTL {
+ e.mu.Unlock()
+ return g.cur, nil
+ }
+ e.mu.Unlock()
+
+ v, err, _ := e.sf.Do(virt.Name, func() (any, error) {
+ repo, err := e.mergeRPM(context.WithoutCancel(ctx), virt)
+ if err != nil {
+ return nil, err
+ }
+ e.mu.Lock()
+ defer e.mu.Unlock()
+ if e.rpm == nil {
+ e.rpm = map[string]*rpmGen{}
+ }
+ g := e.rpm[virt.Name]
+ switch {
+ case g == nil:
+ e.rpm[virt.Name] = &rpmGen{cur: repo, at: time.Now()}
+ case string(g.cur.Repomd) == string(repo.Repomd):
+ g.at = time.Now()
+ repo = g.cur
+ default:
+ g.prev, g.cur, g.at = g.cur, repo, time.Now()
+ }
+ return repo, nil
+ })
+ if err != nil {
+ return nil, err
+ }
+ return v.(*RPMRepo), nil
+}
+
+func (e *Engine) mergeRPM(ctx context.Context, virt models.Virtual) (*RPMRepo, error) {
+ members := make([]RPMMember, len(virt.Members))
+ errs := make([]error, len(virt.Members))
+ var wg sync.WaitGroup
+ for i, name := range virt.Members {
+ wg.Add(1)
+ go func() {
+ defer wg.Done()
+ m, err := e.rpmMember(ctx, name)
+ if err != nil {
+ errs[i] = fmt.Errorf("member %q: %w", name, err)
+ return
+ }
+ members[i] = *m
+ }()
+ }
+ wg.Wait()
+ // %v, not %w: a member failure is a 502 whatever its cause, never a 404.
+ if err := errors.Join(errs...); err != nil {
+ return nil, fmt.Errorf("virtual %q: %v", virt.Name, err)
+ }
+
+ repo, err := MergeRPM(members)
+ if err != nil {
+ return nil, fmt.Errorf("merge rpm repodata: %w", err)
+ }
+ return repo, nil
+}
+
+func (e *Engine) cachedRPMFile(virt, path string) ([]byte, bool) {
+ e.mu.Lock()
+ defer e.mu.Unlock()
+ g := e.rpm[virt]
+ if g == nil {
+ return nil, false
+ }
+ for _, r := range []*RPMRepo{g.cur, g.prev} {
+ if r == nil {
+ continue
+ }
+ if body, ok := r.Files[path]; ok {
+ return body, true
+ }
+ }
+ return nil, false
+}
+
+func (e *Engine) fetchRPMMember(ctx context.Context, name string) (*RPMMember, error) {
+ remote, err := e.getRemote(ctx, name)
+ if err != nil {
+ return nil, fmt.Errorf("remote %q: %w", name, err)
+ }
+ repomd, err := e.fetchMemberPath(ctx, *remote, "repodata/repomd.xml")
+ if err != nil {
+ return nil, err
+ }
+ locs, ts, err := parseRepomd(repomd)
+ if err != nil {
+ return nil, err
+ }
+ m := &RPMMember{RemoteName: name, Bases: remote.UpstreamPool(), Timestamp: ts, Data: map[string][]byte{}}
+ for t, href := range locs {
+ raw, err := e.fetchMemberPath(ctx, *remote, href)
+ if err != nil {
+ return nil, err
+ }
+ if m.Data[t], err = decompress(href, raw); err != nil {
+ return nil, fmt.Errorf("decompress %s: %w", href, err)
+ }
+ }
+ return m, nil
+}
+
+// fetchMemberPath reads one path from a member the same way its own route
+// would serve it: local index, synthesized remote (e.g. github_rpm), or proxy.
+func (e *Engine) fetchMemberPath(ctx context.Context, remote models.Remote, path string) ([]byte, error) {
+ prov, err := provider.Get(remote.PackageType)
+ if err != nil {
+ return nil, fmt.Errorf("provider %q: %w", remote.PackageType, err)
+ }
+
+ var serve func(w http.ResponseWriter, r *http.Request) bool
+ if remote.RepoType == models.RepoTypeLocal {
+ indexer, ok := prov.(provider.LocalIndexer)
+ if !ok {
+ return nil, fmt.Errorf("provider %q does not serve a local index", remote.PackageType)
+ }
+ serve = func(w http.ResponseWriter, r *http.Request) bool {
+ return indexer.ServeLocalIndex(w, r, e.db, remote.Name, path)
+ }
+ } else if rs, ok := prov.(provider.RemoteServer); ok {
+ serve = func(w http.ResponseWriter, r *http.Request) bool {
+ return rs.ServeRemote(w, r, remote, path, "", e.db)
+ }
+ }
+
+ if serve != nil {
+ rec := httptest.NewRecorder()
+ if serve(rec, httptest.NewRequestWithContext(ctx, http.MethodGet, "/"+path, nil)) {
+ if rec.Code != http.StatusOK {
+ return nil, fmt.Errorf("%s/%s: status %d", remote.Name, path, rec.Code)
+ }
+ return rec.Body.Bytes(), nil
+ }
+ if remote.RepoType == models.RepoTypeLocal {
+ return nil, fmt.Errorf("%s/%s: %w", remote.Name, path, ErrNotFound)
+ }
+ }
+
+ res, err := e.proxyEngine.Fetch(ctx, remote, path, prov)
+ if err != nil {
+ return nil, fmt.Errorf("fetch %s/%s: %w", remote.Name, path, err)
+ }
+ defer res.Reader.Close()
+ return io.ReadAll(res.Reader)
+}
diff --git a/internal/virtual/engine_test.go b/internal/virtual/engine_test.go
new file mode 100644
index 0000000..d1c28c9
--- /dev/null
+++ b/internal/virtual/engine_test.go
@@ -0,0 +1,158 @@
+package virtual
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "regexp"
+ "sync"
+ "sync/atomic"
+ "testing"
+ "time"
+
+ "git.unkin.net/unkin/artifactapi/pkg/models"
+)
+
+func fakeEngine(members map[string]*RPMMember) *Engine {
+ return &Engine{
+ getRemote: func(_ context.Context, name string) (*models.Remote, error) {
+ return &models.Remote{Name: name, RepoType: models.RepoTypeRemote}, nil
+ },
+ rpmMember: func(_ context.Context, name string) (*RPMMember, error) {
+ if m := members[name]; m != nil {
+ return m, nil
+ }
+ return nil, errors.New("upstream down")
+ },
+ }
+}
+
+func rpmVirt(members ...string) models.Virtual {
+ return models.Virtual{Name: "v", PackageType: models.PackageRPM, Members: members}
+}
+
+func TestFetchRPMFailsClosed(t *testing.T) {
+ e := fakeEngine(map[string]*RPMMember{
+ "a": {RemoteName: "a", Data: map[string][]byte{"primary": primaryXML(primaryPkgXML("foo", "1", "aaa", "foo.rpm"))}},
+ })
+ _, _, err := e.Fetch(context.Background(), rpmVirt("a", "b"), "repodata/repomd.xml", "")
+ if err == nil || errors.Is(err, ErrNotFound) {
+ t.Fatalf("want upstream error with a member down, got %v", err)
+ }
+}
+
+func TestFetchRPMDataSurvivesMemberChange(t *testing.T) {
+ m := &RPMMember{RemoteName: "a", Data: map[string][]byte{"primary": primaryXML(primaryPkgXML("foo", "1", "aaa", "foo.rpm"))}}
+ e := fakeEngine(map[string]*RPMMember{"a": m})
+ virt := rpmVirt("a")
+
+ repomd, _, err := e.Fetch(context.Background(), virt, "repodata/repomd.xml", "")
+ if err != nil {
+ t.Fatal(err)
+ }
+ m.Data = map[string][]byte{"primary": primaryXML(primaryPkgXML("foo", "2", "bbb", "foo2.rpm"))}
+
+ href := regexp.MustCompile(`repodata/[0-9a-f]+-primary\.xml\.gz`).Find(repomd)
+ body, _, err := e.Fetch(context.Background(), virt, string(href), "")
+ if err != nil {
+ t.Fatalf("data from previous repomd must stay resolvable: %v", err)
+ }
+ if got := string(gunzip(t, body)); !regexp.MustCompile(`ver="1"`).MatchString(got) {
+ t.Fatalf("served wrong primary:\n%s", got)
+ }
+
+ newMD, _, err := e.Fetch(context.Background(), virt, "repodata/repomd.xml", "")
+ if err != nil || string(newMD) == string(repomd) {
+ t.Fatalf("repomd should reflect the member change (err %v)", err)
+ }
+ if _, _, err := e.Fetch(context.Background(), virt, string(href), ""); err != nil {
+ t.Fatalf("previous generation must survive one re-merge: %v", err)
+ }
+}
+
+func TestFetchRPMLocalMemberWithoutRepodataIsUpstreamError(t *testing.T) {
+ e := fakeEngine(map[string]*RPMMember{
+ "a": {RemoteName: "a", Data: map[string][]byte{"primary": primaryXML(primaryPkgXML("foo", "1", "aaa", "foo.rpm"))}},
+ })
+ e.rpmMember = func(_ context.Context, name string) (*RPMMember, error) {
+ if name == "local" {
+ return nil, fmt.Errorf("local/repodata/repomd.xml: %w", ErrNotFound)
+ }
+ return &RPMMember{RemoteName: name, Data: map[string][]byte{"primary": primaryXML(primaryPkgXML("foo", "1", "aaa", "foo.rpm"))}}, nil
+ }
+ _, _, err := e.Fetch(context.Background(), rpmVirt("a", "local"), "repodata/repomd.xml", "")
+ if err == nil || errors.Is(err, ErrNotFound) {
+ t.Fatalf("member failure must not surface as not-found, got %v", err)
+ }
+}
+
+func TestFetchRPMMergeIsShared(t *testing.T) {
+ var calls atomic.Int32
+ e := fakeEngine(nil)
+ e.mergeTTL = time.Minute
+ e.rpmMember = func(_ context.Context, name string) (*RPMMember, error) {
+ calls.Add(1)
+ time.Sleep(20 * time.Millisecond)
+ return &RPMMember{RemoteName: name, Data: map[string][]byte{"primary": primaryXML(primaryPkgXML("foo", "1", "aaa", "foo.rpm"))}}, nil
+ }
+ var wg sync.WaitGroup
+ for range 10 {
+ wg.Add(1)
+ go func() {
+ defer wg.Done()
+ if _, _, err := e.Fetch(context.Background(), rpmVirt("a"), "repodata/repomd.xml", ""); err != nil {
+ t.Error(err)
+ }
+ }()
+ }
+ wg.Wait()
+ if _, _, err := e.Fetch(context.Background(), rpmVirt("a"), "repodata/repomd.xml", ""); err != nil {
+ t.Fatal(err)
+ }
+ if n := calls.Load(); n != 1 {
+ t.Fatalf("member fetched %d times, want 1", n)
+ }
+}
+
+func TestFetchRPMDataScopedToVirtual(t *testing.T) {
+ e := fakeEngine(map[string]*RPMMember{
+ "a": {RemoteName: "a", Data: map[string][]byte{"primary": primaryXML(primaryPkgXML("foo", "1", "aaa", "foo.rpm"))}},
+ "b": {RemoteName: "b", Data: map[string][]byte{"primary": primaryXML(primaryPkgXML("bar", "1", "bbb", "bar.rpm"))}},
+ })
+ virtA := models.Virtual{Name: "va", PackageType: models.PackageRPM, Members: []string{"a"}}
+ virtB := models.Virtual{Name: "vb", PackageType: models.PackageRPM, Members: []string{"b"}}
+
+ repomd, _, err := e.Fetch(context.Background(), virtA, "repodata/repomd.xml", "")
+ if err != nil {
+ t.Fatal(err)
+ }
+ href := regexp.MustCompile(`repodata/[0-9a-f]+-primary\.xml\.gz`).Find(repomd)
+ if _, _, err := e.Fetch(context.Background(), virtB, string(href), ""); !errors.Is(err, ErrNotFound) {
+ t.Fatalf("virtual vb served va's data file (err %v)", err)
+ }
+}
+
+func TestMemberRedirect(t *testing.T) {
+ e := fakeEngine(nil)
+ virt := rpmVirt("gh", "other")
+ for _, tc := range []struct {
+ path, query, want string
+ ok bool
+ err error
+ }{
+ {path: "gh/o/r/releases/download/v1/c++-1.rpm", query: "x=1", ok: true,
+ want: "http://h/api/v1/remote/gh/o/r/releases/download/v1/c++-1.rpm?x=1"},
+ {path: "gh/a%20b/c%3Fd.rpm", ok: true, want: "http://h/api/v1/remote/gh/a%20b/c%3Fd.rpm"},
+ {path: "gh/../../remote/other/x", err: ErrBadPath},
+ {path: "gh/%2e%2e/%2e%2e/remote/other/x", err: ErrBadPath},
+ {path: "gh/a/./b", err: ErrBadPath},
+ {path: "gh//etc/passwd", err: ErrBadPath},
+ {path: "repodata/repomd.xml"},
+ {path: "nope/x.rpm"},
+ } {
+ got, ok, err := e.MemberRedirect(context.Background(), virt, tc.path, tc.query, "http://h/")
+ if got != tc.want || ok != tc.ok || !errors.Is(err, tc.err) {
+ t.Errorf("%s: got (%q, %v, %v), want (%q, %v, %v)", tc.path, got, ok, err, tc.want, tc.ok, tc.err)
+ }
+ }
+}
diff --git a/internal/virtual/rpm_merger.go b/internal/virtual/rpm_merger.go
new file mode 100644
index 0000000..53f5c50
--- /dev/null
+++ b/internal/virtual/rpm_merger.go
@@ -0,0 +1,298 @@
+package virtual
+
+import (
+ "bytes"
+ "compress/bzip2"
+ "compress/gzip"
+ "crypto/sha256"
+ "encoding/hex"
+ "encoding/xml"
+ "errors"
+ "fmt"
+ "io"
+ "regexp"
+ "strings"
+
+ "github.com/klauspost/compress/zstd"
+ "github.com/ulikunitz/xz"
+)
+
+var rpmDataTypes = []string{"primary", "filelists", "other"}
+
+var rpmRoots = map[string]string{
+ "primary": ``,
+ "filelists": ``,
+ "other": ``,
+}
+
+var rpmRootClose = map[string]string{"primary": "", "filelists": "", "other": ""}
+
+var locationRe = regexp.MustCompile(`]*/>`)
+
+// RPMMember is one member repo's decompressed primary/filelists/other XML,
+// keyed by repomd data type. Missing types contribute no packages.
+type RPMMember struct {
+ RemoteName string
+ Bases []string
+ Timestamp int64
+ Data map[string][]byte
+}
+
+// RPMRepo is a merged yum repo: repomd.xml plus the files it references, keyed
+// by their path under the repo root.
+type RPMRepo struct {
+ Repomd []byte
+ Files map[string][]byte
+}
+
+type repomdDoc struct {
+ Revision string `xml:"revision"`
+ Data []repomdData `xml:"data"`
+}
+
+type repomdData struct {
+ Type string `xml:"type,attr"`
+ Location struct {
+ Href string `xml:"href,attr"`
+ } `xml:"location"`
+ Timestamp int64 `xml:"timestamp"`
+}
+
+type rpmPkgVersion struct {
+ Epoch string `xml:"epoch,attr"`
+ Ver string `xml:"ver,attr"`
+ Rel string `xml:"rel,attr"`
+}
+
+type primaryPkg struct {
+ Name string `xml:"name"`
+ Arch string `xml:"arch"`
+ Version rpmPkgVersion `xml:"version"`
+ Checksum string `xml:"checksum"`
+ Location struct {
+ Href string `xml:"href,attr"`
+ Base string `xml:"http://www.w3.org/XML/1998/namespace base,attr"`
+ } `xml:"location"`
+}
+
+type pkgidPkg struct {
+ PkgID string `xml:"pkgid,attr"`
+}
+
+// parseRepomd returns the member's primary/filelists/other locations and its
+// newest data timestamp.
+func parseRepomd(body []byte) (map[string]string, int64, error) {
+ var doc repomdDoc
+ if err := xml.Unmarshal(body, &doc); err != nil {
+ return nil, 0, fmt.Errorf("parse repomd.xml: %w", err)
+ }
+ locs := map[string]string{}
+ var ts int64
+ for _, d := range doc.Data {
+ for _, t := range rpmDataTypes {
+ if d.Type == t && d.Location.Href != "" {
+ locs[t] = d.Location.Href
+ ts = max(ts, d.Timestamp)
+ }
+ }
+ }
+ if locs["primary"] == "" {
+ return nil, 0, errors.New("repomd.xml has no primary data")
+ }
+ return locs, ts, nil
+}
+
+func decompress(href string, body []byte) ([]byte, error) {
+ var r io.Reader
+ switch {
+ case strings.HasSuffix(href, ".gz"):
+ gz, err := gzip.NewReader(bytes.NewReader(body))
+ if err != nil {
+ return nil, err
+ }
+ r = gz
+ case strings.HasSuffix(href, ".xz"):
+ x, err := xz.NewReader(bytes.NewReader(body))
+ if err != nil {
+ return nil, err
+ }
+ r = x
+ case strings.HasSuffix(href, ".zst"):
+ z, err := zstd.NewReader(bytes.NewReader(body))
+ if err != nil {
+ return nil, err
+ }
+ defer z.Close()
+ r = z
+ case strings.HasSuffix(href, ".bz2"):
+ r = bzip2.NewReader(bytes.NewReader(body))
+ default:
+ return body, nil
+ }
+ return io.ReadAll(r)
+}
+
+// splitPackages returns the raw bytes of each top-level element.
+func splitPackages(doc []byte) ([][]byte, error) {
+ d := xml.NewDecoder(bytes.NewReader(doc))
+ var pkgs [][]byte
+ depth := 0
+ for {
+ start := d.InputOffset()
+ tok, err := d.Token()
+ if err == io.EOF {
+ return pkgs, nil
+ }
+ if err != nil {
+ return nil, err
+ }
+ switch t := tok.(type) {
+ case xml.StartElement:
+ if depth == 1 && t.Name.Local == "package" {
+ if err := d.Skip(); err != nil {
+ return nil, err
+ }
+ pkgs = append(pkgs, doc[start:d.InputOffset()])
+ continue
+ }
+ depth++
+ case xml.EndElement:
+ depth--
+ }
+ }
+}
+
+// memberHref returns href relative to the member root. An xml:base under one of
+// the member's upstream bases is folded into the href; any other xml:base is a
+// host the member can't serve, so ok is false and the location is left as is.
+func memberHref(m RPMMember, base, href string) (string, bool) {
+ href = strings.TrimLeft(href, "/")
+ if base == "" {
+ return href, true
+ }
+ base = strings.TrimRight(base, "/") + "/"
+ for _, b := range m.Bases {
+ b = strings.TrimRight(b, "/") + "/"
+ if strings.HasPrefix(base, b) {
+ return strings.TrimPrefix(base, b) + href, true
+ }
+ }
+ return "", false
+}
+
+func xmlEscape(s string) string {
+ var b bytes.Buffer
+ _ = xml.EscapeText(&b, []byte(s))
+ return b.String()
+}
+
+// MergeRPM merges members into one repo. Members are in priority order: the
+// first member to carry a NEVRA wins. Package locations are prefixed with the
+// owning member's name so the virtual can route downloads back to it.
+func MergeRPM(members []RPMMember) (*RPMRepo, error) {
+ kept := map[string]bool{}
+ seen := map[string]bool{}
+ out := map[string][][]byte{}
+
+ for i, m := range members {
+ pkgs, err := splitPackages(m.Data["primary"])
+ if err != nil {
+ return nil, fmt.Errorf("member %q primary: %w", m.RemoteName, err)
+ }
+ for _, raw := range pkgs {
+ var p primaryPkg
+ if err := xml.Unmarshal(raw, &p); err != nil {
+ return nil, fmt.Errorf("member %q primary package: %w", m.RemoteName, err)
+ }
+ epoch := p.Version.Epoch
+ if epoch == "" {
+ epoch = "0"
+ }
+ nevra := fmt.Sprintf("%s-%s:%s-%s.%s", p.Name, epoch, p.Version.Ver, p.Version.Rel, p.Arch)
+ if seen[nevra] {
+ continue
+ }
+ seen[nevra] = true
+ kept[fmt.Sprintf("%d/%s", i, p.Checksum)] = true
+
+ if href, ok := memberHref(m, p.Location.Base, p.Location.Href); ok {
+ loc := ``
+ raw = locationRe.ReplaceAll(raw, []byte(loc))
+ }
+ out["primary"] = append(out["primary"], raw)
+ }
+ }
+
+ for _, t := range rpmDataTypes[1:] {
+ emitted := map[string]bool{}
+ for i, m := range members {
+ pkgs, err := splitPackages(m.Data[t])
+ if err != nil {
+ return nil, fmt.Errorf("member %q %s: %w", m.RemoteName, t, err)
+ }
+ for _, raw := range pkgs {
+ var p pkgidPkg
+ if err := xml.Unmarshal(raw, &p); err != nil {
+ return nil, fmt.Errorf("member %q %s package: %w", m.RemoteName, t, err)
+ }
+ key := fmt.Sprintf("%d/%s", i, p.PkgID)
+ if kept[key] && !emitted[key] {
+ emitted[key] = true
+ out[t] = append(out[t], raw)
+ }
+ }
+ }
+ }
+
+ var ts int64
+ for _, m := range members {
+ ts = max(ts, m.Timestamp)
+ }
+
+ repo := &RPMRepo{Files: map[string][]byte{}}
+ var md bytes.Buffer
+ md.WriteString(xml.Header)
+ md.WriteString(`` + "\n")
+ fmt.Fprintf(&md, " %d\n", ts)
+ for _, t := range rpmDataTypes {
+ var doc bytes.Buffer
+ doc.WriteString(xml.Header)
+ fmt.Fprintf(&doc, rpmRoots[t]+"\n", len(out[t]))
+ for _, raw := range out[t] {
+ doc.Write(raw)
+ doc.WriteString("\n")
+ }
+ doc.WriteString(rpmRootClose[t] + "\n")
+
+ gz := gzipDeterministic(doc.Bytes())
+ sum := sha256Hex(gz)
+ href := fmt.Sprintf("repodata/%s-%s.xml.gz", sum, t)
+ repo.Files[href] = gz
+
+ fmt.Fprintf(&md, " \n", t)
+ fmt.Fprintf(&md, " %s\n", sum)
+ fmt.Fprintf(&md, " %s\n", sha256Hex(doc.Bytes()))
+ fmt.Fprintf(&md, " \n", href)
+ fmt.Fprintf(&md, " %d\n", ts)
+ fmt.Fprintf(&md, " %d\n", len(gz))
+ fmt.Fprintf(&md, " %d\n", doc.Len())
+ md.WriteString(" \n")
+ }
+ md.WriteString("\n")
+ repo.Repomd = md.Bytes()
+ return repo, nil
+}
+
+func gzipDeterministic(data []byte) []byte {
+ var buf bytes.Buffer
+ gz := gzip.NewWriter(&buf)
+ gz.Header = gzip.Header{OS: 255}
+ _, _ = gz.Write(data)
+ _ = gz.Close()
+ return buf.Bytes()
+}
+
+func sha256Hex(data []byte) string {
+ h := sha256.Sum256(data)
+ return hex.EncodeToString(h[:])
+}
diff --git a/internal/virtual/rpm_merger_test.go b/internal/virtual/rpm_merger_test.go
new file mode 100644
index 0000000..cda6d50
--- /dev/null
+++ b/internal/virtual/rpm_merger_test.go
@@ -0,0 +1,274 @@
+package virtual
+
+import (
+ "bytes"
+ "compress/gzip"
+ "crypto/sha256"
+ "encoding/hex"
+ "encoding/xml"
+ "fmt"
+ "io"
+ "strings"
+ "testing"
+
+ "github.com/klauspost/compress/zstd"
+ "github.com/ulikunitz/xz"
+)
+
+func primaryXML(pkgs ...string) []byte {
+ return []byte(`
+
+` + strings.Join(pkgs, "\n") + `
+`)
+}
+
+func primaryPkgXML(name, ver, pkgid, href string) string {
+ return `
+ ` + name + `
+ x86_64
+
+ ` + pkgid + `
+
+
+
+
+`
+}
+
+func filelistsXML(pkgids ...string) []byte {
+ var b strings.Builder
+ b.WriteString(``)
+ for _, id := range pkgids {
+ b.WriteString(`/usr/bin/` + id + ``)
+ }
+ b.WriteString(``)
+ return []byte(b.String())
+}
+
+func gunzip(t *testing.T, b []byte) []byte {
+ t.Helper()
+ r, err := gzip.NewReader(bytes.NewReader(b))
+ if err != nil {
+ t.Fatal(err)
+ }
+ out, err := io.ReadAll(r)
+ if err != nil {
+ t.Fatal(err)
+ }
+ return out
+}
+
+func mergedFile(t *testing.T, repo *RPMRepo, dtype string) []byte {
+ t.Helper()
+ for href, body := range repo.Files {
+ if strings.HasSuffix(href, "-"+dtype+".xml.gz") {
+ return gunzip(t, body)
+ }
+ }
+ t.Fatalf("no %s file in merged repo", dtype)
+ return nil
+}
+
+func TestMergeRPMDedupePriority(t *testing.T) {
+ members := []RPMMember{
+ {RemoteName: "first", Timestamp: 100, Data: map[string][]byte{
+ "primary": primaryXML(primaryPkgXML("foo", "1.0", "aaa", "Packages/foo-1.0.rpm")),
+ "filelists": filelistsXML("aaa"),
+ }},
+ {RemoteName: "second", Timestamp: 200, Data: map[string][]byte{
+ "primary": primaryXML(
+ primaryPkgXML("foo", "1.0", "bbb", "Packages/foo-1.0-rebuilt.rpm"),
+ primaryPkgXML("bar", "2.0", "ccc", "Packages/bar-2.0.rpm"),
+ ),
+ "filelists": filelistsXML("bbb", "ccc"),
+ }},
+ }
+ repo, err := MergeRPM(members)
+ if err != nil {
+ t.Fatal(err)
+ }
+
+ primary := string(mergedFile(t, repo, "primary"))
+ if strings.Count(primary, "foo") != 1 {
+ t.Fatalf("duplicate NEVRA not deduped:\n%s", primary)
+ }
+ if !strings.Contains(primary, `href="first/Packages/foo-1.0.rpm"`) || strings.Contains(primary, "foo-1.0-rebuilt") {
+ t.Fatalf("first member should win duplicate NEVRA:\n%s", primary)
+ }
+ if !strings.Contains(primary, `href="second/Packages/bar-2.0.rpm"`) {
+ t.Fatalf("unique package from second member missing:\n%s", primary)
+ }
+ if !strings.Contains(primary, `packages="2"`) {
+ t.Fatalf("package count wrong:\n%s", primary)
+ }
+
+ filelists := string(mergedFile(t, repo, "filelists"))
+ if !strings.Contains(filelists, "/usr/bin/aaa") || !strings.Contains(filelists, "/usr/bin/ccc") || strings.Contains(filelists, "/usr/bin/bbb") {
+ t.Fatalf("filelists must follow the winning primary entries:\n%s", filelists)
+ }
+
+ other := string(mergedFile(t, repo, "other"))
+ if !strings.Contains(other, `packages="0"`) {
+ t.Fatalf("members without other data should yield an empty other.xml:\n%s", other)
+ }
+
+ if !strings.Contains(string(repo.Repomd), "200") {
+ t.Fatalf("revision should be the newest member timestamp:\n%s", repo.Repomd)
+ }
+}
+
+func TestMergeRPMOutputIsWellFormed(t *testing.T) {
+ repo, err := MergeRPM([]RPMMember{{RemoteName: "a", Data: map[string][]byte{
+ "primary": primaryXML(primaryPkgXML("foo", "1.0", "aaa", "Packages/foo.rpm")),
+ }}})
+ if err != nil {
+ t.Fatal(err)
+ }
+ var doc struct {
+ XMLName xml.Name
+ Packages []struct {
+ Name string `xml:"name"`
+ Provides []struct {
+ Name string `xml:"name,attr"`
+ } `xml:"format>provides>entry"`
+ } `xml:"package"`
+ }
+ if err := xml.Unmarshal(mergedFile(t, repo, "primary"), &doc); err != nil {
+ t.Fatal(err)
+ }
+ if doc.XMLName.Space != "http://linux.duke.edu/metadata/common" || len(doc.Packages) != 1 {
+ t.Fatalf("unexpected merged primary: %+v", doc)
+ }
+ if len(doc.Packages[0].Provides) != 1 || doc.Packages[0].Provides[0].Name != "foo" {
+ t.Fatalf("rpm: namespaced format lost: %+v", doc.Packages[0])
+ }
+}
+
+func TestMergeRPMHrefRewriting(t *testing.T) {
+ based := func(name, base string) string {
+ return `` + name + `noarch` +
+ `` + name + ``
+ }
+ repo, err := MergeRPM([]RPMMember{{RemoteName: "gh", Bases: []string{"https://up.example/el9", "https://mirror.example/el9/"}, Data: map[string][]byte{
+ "primary": primaryXML(
+ primaryPkgXML("foo", "1.0", "aaa", "storytold/photocraft/releases/download/v1.0/foo&bar.rpm"),
+ primaryPkgXML("lead", "1.0", "bbb", "/Packages/lead.rpm"),
+ based("root", "https://up.example/el9/"),
+ based("sub", "https://mirror.example/el9/extra"),
+ based("ext", "https://other.example/"),
+ ),
+ }}})
+ if err != nil {
+ t.Fatal(err)
+ }
+ primary := string(mergedFile(t, repo, "primary"))
+ for _, want := range []string{
+ ``,
+ ``,
+ ``,
+ ``,
+ ``,
+ } {
+ if !strings.Contains(primary, want) {
+ t.Errorf("missing %s in:\n%s", want, primary)
+ }
+ }
+}
+
+func TestMergeRPMChecksums(t *testing.T) {
+ repo, err := MergeRPM([]RPMMember{{RemoteName: "a", Timestamp: 42, Data: map[string][]byte{
+ "primary": primaryXML(primaryPkgXML("foo", "1.0", "aaa", "Packages/foo.rpm")),
+ }}})
+ if err != nil {
+ t.Fatal(err)
+ }
+
+ var md struct {
+ Data []struct {
+ Type string `xml:"type,attr"`
+ Checksum string `xml:"checksum"`
+ OpenChecksum string `xml:"open-checksum"`
+ Location struct {
+ Href string `xml:"href,attr"`
+ } `xml:"location"`
+ Timestamp int64 `xml:"timestamp"`
+ Size int `xml:"size"`
+ OpenSize int `xml:"open-size"`
+ } `xml:"data"`
+ }
+ if err := xml.Unmarshal(repo.Repomd, &md); err != nil {
+ t.Fatal(err)
+ }
+ if len(md.Data) != 3 {
+ t.Fatalf("want primary/filelists/other, got %d entries", len(md.Data))
+ }
+ for _, d := range md.Data {
+ body, ok := repo.Files[d.Location.Href]
+ if !ok {
+ t.Fatalf("%s: location %q not served", d.Type, d.Location.Href)
+ }
+ sum := sha256.Sum256(body)
+ if hex.EncodeToString(sum[:]) != d.Checksum || len(body) != d.Size {
+ t.Errorf("%s: checksum/size do not match served bytes", d.Type)
+ }
+ open := gunzip(t, body)
+ osum := sha256.Sum256(open)
+ if hex.EncodeToString(osum[:]) != d.OpenChecksum || len(open) != d.OpenSize {
+ t.Errorf("%s: open-checksum/open-size do not match decompressed bytes", d.Type)
+ }
+ if d.Timestamp != 42 {
+ t.Errorf("%s: timestamp %d, want 42", d.Type, d.Timestamp)
+ }
+ }
+
+ again, _ := MergeRPM([]RPMMember{{RemoteName: "a", Timestamp: 42, Data: map[string][]byte{
+ "primary": primaryXML(primaryPkgXML("foo", "1.0", "aaa", "Packages/foo.rpm")),
+ }}})
+ if !bytes.Equal(repo.Repomd, again.Repomd) {
+ t.Error("merge must be deterministic so repomd checksums stay valid across requests")
+ }
+}
+
+func TestParseRepomd(t *testing.T) {
+ locs, ts, err := parseRepomd([]byte(`
+ 10
+ 99
+ 20
+`))
+ if err != nil {
+ t.Fatal(err)
+ }
+ if locs["primary"] != "repodata/p-primary.xml.zst" || locs["other"] != "repodata/o-other.xml.gz" || len(locs) != 2 {
+ t.Fatalf("unexpected locations: %v", locs)
+ }
+ if ts != 20 {
+ t.Fatalf("timestamp %d, want 20 (sqlite entries ignored)", ts)
+ }
+ if _, _, err := parseRepomd([]byte(``)); err == nil {
+ t.Fatal("repomd without primary should error")
+ }
+}
+
+func TestDecompress(t *testing.T) {
+ plain := []byte("")
+
+ var xzBuf bytes.Buffer
+ xw, _ := xz.NewWriter(&xzBuf)
+ _, _ = xw.Write(plain)
+ _ = xw.Close()
+
+ zw, _ := zstd.NewWriter(nil)
+ zst := zw.EncodeAll(plain, nil)
+
+ for href, body := range map[string][]byte{
+ "p.xml": plain,
+ "p.xml.gz": gzipDeterministic(plain),
+ "p.xml.xz": xzBuf.Bytes(),
+ "p.xml.zst": zst,
+ } {
+ got, err := decompress(href, body)
+ if err != nil || !bytes.Equal(got, plain) {
+ t.Errorf("%s: got %q, %v", href, got, err)
+ }
+ }
+}