Files
artifactapi/internal/virtual/engine.go
T
unkin-agent 1837f6ef8c
ci/woodpecker/tag/docker Pipeline failed
Add rpm virtual repositories (#131)
Virtual repos only merge helm and pypi, so several rpm repos (e.g. many github_rpm remotes) cannot be served as one yum repo.

- merge member primary/filelists/other into one repodata set; member order wins duplicate NEVRAs
- read member repodata through local, github_rpm and proxied remote paths
- prefix package locations (incl. xml:base under a member upstream) with the member name and 302 them to the member route
- reject absolute and dot-segment member paths; escape the redirect and keep its query
- return 502 when any member's repodata is unavailable
- reuse each virtual's merge for 60s behind singleflight; serve data files from its current and previous merge (per replica)
- add merger/engine unit tests and a dockerised dnf e2e case

Reviewed-on: #131
Co-authored-by: unkin-agent <unkin-agent@unkin.net>
Co-committed-by: unkin-agent <unkin-agent@unkin.net>
2026-10-09 22:07:20 +11:00

397 lines
11 KiB
Go

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 {
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)
}
members, err := e.fetchMemberIndexes(ctx, virt, path)
if err != nil {
return nil, "", err
}
if len(members) == 0 {
return nil, "", fmt.Errorf("no members reachable for virtual %q", virt.Name)
}
merged, err := merger.MergeIndexes(members, proxyBaseURL)
if err != nil {
return nil, "", fmt.Errorf("merge indexes: %w", err)
}
contentType := "application/octet-stream"
switch virt.PackageType {
case models.PackageHelm:
contentType = "text/yaml"
case models.PackagePyPI:
contentType = "text/html"
}
return merged, contentType, nil
}
func (e *Engine) fetchMemberIndexes(ctx context.Context, virt models.Virtual, path string) ([]MemberIndex, error) {
type result struct {
index MemberIndex
err error
}
results := make([]result, len(virt.Members))
var wg sync.WaitGroup
for i, memberName := range virt.Members {
wg.Add(1)
go func(idx int, name string) {
defer wg.Done()
remote, err := e.db.GetRemote(ctx, name)
if err != nil {
results[idx] = result{err: fmt.Errorf("remote %q: %w", name, err)}
return
}
if remote.RepoType == models.RepoTypeLocal {
body, err := e.fetchLocalIndex(ctx, *remote, path)
if err != nil {
results[idx] = result{err: fmt.Errorf("local index %q: %w", name, err)}
return
}
results[idx] = result{index: MemberIndex{RemoteName: name, RepoType: remote.RepoType, BaseURL: remote.BaseURL, Body: body}}
return
}
prov, err := provider.Get(remote.PackageType)
if err != nil {
results[idx] = result{err: fmt.Errorf("provider %q: %w", remote.PackageType, err)}
return
}
fetchResult, err := e.proxyEngine.Fetch(ctx, *remote, path, prov)
if err != nil {
results[idx] = result{err: fmt.Errorf("fetch %q/%s: %w", name, path, err)}
return
}
defer fetchResult.Reader.Close()
body, err := io.ReadAll(fetchResult.Reader)
if err != nil {
results[idx] = result{err: fmt.Errorf("read %q: %w", name, err)}
return
}
results[idx] = result{index: MemberIndex{RemoteName: name, RepoType: remote.RepoType, BaseURL: remote.BaseURL, Body: body}}
}(i, memberName)
}
wg.Wait()
var members []MemberIndex
for _, r := range results {
if r.err != nil {
slog.Warn("virtual member fetch failed", "error", r.err)
continue
}
members = append(members, r.index)
}
return members, nil
}
func (e *Engine) fetchLocalIndex(ctx context.Context, remote models.Remote, path string) ([]byte, error) {
prov, err := provider.Get(remote.PackageType)
if err != nil {
return nil, fmt.Errorf("no provider for %q: %w", remote.PackageType, err)
}
indexer, ok := prov.(provider.LocalIndexer)
if !ok {
return nil, fmt.Errorf("provider %q does not support local index generation", remote.PackageType)
}
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 (<member>/<href>, 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)
}