589421f995
- reject absolute and dot-segment member paths, escape the target, keep the query - return 502 when any member's repodata can't be fetched - serve recently merged data files by hash for 5m after repomd changes - fold xml:base under a member upstream into the virtual href
355 lines
9.8 KiB
Go
355 lines
9.8 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"
|
|
)
|
|
|
|
type Engine struct {
|
|
db *database.DB
|
|
proxyEngine *proxy.Engine
|
|
getRemote func(context.Context, string) (*models.Remote, error)
|
|
rpmMember func(context.Context, string) (*RPMMember, error)
|
|
|
|
mu sync.Mutex
|
|
rpmFiles map[string]rpmFile
|
|
}
|
|
|
|
type rpmFile struct {
|
|
body []byte
|
|
expires time.Time
|
|
}
|
|
|
|
// rpmFileTTL keeps content-hashed data files resolvable after repomd moves on,
|
|
// so a client that fetched the previous repomd can still fetch its data.
|
|
const rpmFileTTL = 5 * time.Minute
|
|
|
|
func NewEngine(db *database.DB, proxyEngine *proxy.Engine) *Engine {
|
|
e := &Engine{db: db, proxyEngine: proxyEngine, getRemote: db.GetRemote}
|
|
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(path); ok {
|
|
return body, "application/gzip", nil
|
|
}
|
|
}
|
|
|
|
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()
|
|
if err := errors.Join(errs...); err != nil {
|
|
return nil, "", fmt.Errorf("virtual %q: %w", virt.Name, err)
|
|
}
|
|
|
|
repo, err := MergeRPM(members)
|
|
if err != nil {
|
|
return nil, "", fmt.Errorf("merge rpm repodata: %w", err)
|
|
}
|
|
e.storeRPMFiles(repo.Files)
|
|
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
|
|
}
|
|
|
|
// 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) storeRPMFiles(files map[string][]byte) {
|
|
now := time.Now()
|
|
e.mu.Lock()
|
|
defer e.mu.Unlock()
|
|
if e.rpmFiles == nil {
|
|
e.rpmFiles = map[string]rpmFile{}
|
|
}
|
|
for k, f := range e.rpmFiles {
|
|
if now.After(f.expires) {
|
|
delete(e.rpmFiles, k)
|
|
}
|
|
}
|
|
for k, body := range files {
|
|
e.rpmFiles[k] = rpmFile{body: body, expires: now.Add(rpmFileTTL)}
|
|
}
|
|
}
|
|
|
|
func (e *Engine) cachedRPMFile(path string) ([]byte, bool) {
|
|
e.mu.Lock()
|
|
defer e.mu.Unlock()
|
|
f, ok := e.rpmFiles[path]
|
|
if !ok || time.Now().After(f.expires) {
|
|
return nil, false
|
|
}
|
|
return f.body, true
|
|
}
|
|
|
|
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)
|
|
}
|