Initial release: self-hostable APT repository server and CLI
urapt is a self-hostable APT repository server with a companion CLI for pushing and managing Debian .deb packages. Server (urapt-server): - REST API + APT endpoint, SQLite storage (pure-Go modernc driver, no CGO) - .deb files stored content-addressed on disk, reference-counted for dedup - Server-managed RSA-4096 OpenPGP signing key (ProtonMail/go-crypto) - APT indices (Release/InRelease/Packages[.gz/.xz]) generated on demand from the DB, cached in memory, signed with the server key - Full APT model: repositories -> distributions -> components -> architectures - Bearer-token auth for REST; HTTP Basic auth for private-repo APT reads - First registrant becomes admin; repo-scoped permissions (read/write/read-write/admin) plus owner and server-admin roles - Multipart package push with control-field extraction, list/show/delete, pool serving, blob ref-count cleanup - Audit log CLI (urapt): - register/login/logout/whoami, token management - repo/distro/component/arch CRUD, member management - push/pull/ls/show/rm for packages - apt-config helper that emits apt setup commands (key, sources.list, auth.conf for private repos) Packaging & docs: - Dockerfile (multi-stage distroless), docker-compose.yml, sample config - README quick start, architecture overview, config reference, security notes - PLAN.md design blueprint, CHANGELOG.md, GPL-3.0 LICENSE - GitHub Actions CI (test, lint, cross-build for linux/darwin amd64/arm64) - Makefile release target producing static binaries + tarballs + checksums Tests cover the data-access layer, auth/permission checks, APT index generation, .deb parsing, GPG signing, the REST API, and the typed API client. Verified end-to-end on a Raspberry Pi (arm64) pushing and installing a real package.
This commit is contained in:
@@ -0,0 +1,211 @@
|
||||
// Package app wires the urapt-server dependencies and runs the HTTP server.
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"flag"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
|
||||
"urapt/server/aptrepo"
|
||||
"urapt/server/auth"
|
||||
"urapt/server/cache"
|
||||
"urapt/server/restapi"
|
||||
"urapt/server/store"
|
||||
"urapt/shared/config"
|
||||
"urapt/shared/db"
|
||||
"urapt/shared/gpg"
|
||||
"urapt/shared/log"
|
||||
)
|
||||
|
||||
// App holds the server configuration and runtime dependencies.
|
||||
type App struct {
|
||||
args []string
|
||||
version string
|
||||
}
|
||||
|
||||
// New constructs an App from the given command-line args and version string.
|
||||
func New(args []string, version string) *App {
|
||||
return &App{args: args, version: version}
|
||||
}
|
||||
|
||||
// signerProvider implements restapi.SignerProvider backed by a *gpg.Key.
|
||||
type signerProvider struct {
|
||||
key *gpg.Key
|
||||
}
|
||||
|
||||
func (s *signerProvider) PublicKeyArmored() (string, error) {
|
||||
if s.key == nil {
|
||||
return "", fmt.Errorf("no key")
|
||||
}
|
||||
return s.key.ArmoredPublic()
|
||||
}
|
||||
|
||||
func (s *signerProvider) Fingerprint() string {
|
||||
if s.key == nil {
|
||||
return ""
|
||||
}
|
||||
return s.key.Fingerprint
|
||||
}
|
||||
|
||||
// Key returns the signing key for the APT endpoint (Phase 6).
|
||||
func (s *signerProvider) Key() *gpg.Key { return s.key }
|
||||
|
||||
// Run loads config, opens the database, ensures a signing key, and serves
|
||||
// HTTP until the context is cancelled.
|
||||
func (a *App) Run(ctx context.Context) error {
|
||||
cfg, err := a.loadConfig()
|
||||
if err != nil {
|
||||
fmt.Fprintln(os.Stderr, "config error:", err)
|
||||
return err
|
||||
}
|
||||
logger := log.New(cfg.LogLevel)
|
||||
logger.Info("starting urapt-server", "version", a.version, "bind", cfg.Bind, "base_url", cfg.BaseURL)
|
||||
|
||||
if err := cfg.Validate(); err != nil {
|
||||
logger.Error("invalid config", "err", err)
|
||||
return err
|
||||
}
|
||||
|
||||
if err := ensureDirs(cfg); err != nil {
|
||||
logger.Error("setup store dirs", "err", err)
|
||||
return err
|
||||
}
|
||||
|
||||
database, err := db.Open(cfg.DBPath)
|
||||
if err != nil {
|
||||
logger.Error("open database", "err", err)
|
||||
return err
|
||||
}
|
||||
defer database.Close()
|
||||
|
||||
st := store.New(database)
|
||||
authSvc := auth.NewService(st)
|
||||
|
||||
signer, err := ensureSigningKey(ctx, st, cfg)
|
||||
if err != nil {
|
||||
logger.Error("ensure signing key", "err", err)
|
||||
return err
|
||||
}
|
||||
sp := &signerProvider{key: signer}
|
||||
logger.Info("signing key ready", "fingerprint", sp.Fingerprint())
|
||||
|
||||
idxCache := cache.New()
|
||||
|
||||
root := chi.NewRouter()
|
||||
root.Mount("/api/v1", restapi.New(st, authSvc, sp, &cfg, idxCache))
|
||||
root.Mount("/apt", aptrepo.New(st, authSvc, sp.Key(), idxCache))
|
||||
|
||||
srv := &http.Server{
|
||||
Addr: cfg.Bind,
|
||||
Handler: root,
|
||||
}
|
||||
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
logger.Info("listening", "addr", cfg.Bind)
|
||||
if cfg.TLSEnabled {
|
||||
errCh <- srv.ListenAndServeTLS(cfg.TLSCert, cfg.TLSKey)
|
||||
} else {
|
||||
errCh <- srv.ListenAndServe()
|
||||
}
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
logger.Info("shutting down")
|
||||
shutCtx, cancel := context.WithTimeout(context.Background(), shutdownTimeout)
|
||||
defer cancel()
|
||||
return srv.Shutdown(shutCtx)
|
||||
case err := <-errCh:
|
||||
if err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||
logger.Error("server error", "err", err)
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// loadConfig parses flags and builds the effective Config.
|
||||
func (a *App) loadConfig() (config.Config, error) {
|
||||
fs := flag.NewFlagSet("urapt-server", flag.ContinueOnError)
|
||||
fs.SetOutput(os.Stderr)
|
||||
configPath := fs.String("config", config.Defaults.ConfigPath, "path to config file")
|
||||
fs.String("bind", "", "bind address")
|
||||
fs.String("base-url", "", "external base URL")
|
||||
fs.String("store-dir", "", "store directory")
|
||||
fs.String("db-path", "", "sqlite db path")
|
||||
fs.String("packages-dir", "", "packages directory")
|
||||
fs.String("log-level", "", "log level")
|
||||
fs.String("signing-key-type", "", "signing key type")
|
||||
fs.Int("signing-key-bits", 0, "signing key bits")
|
||||
fs.String("signing-key-user-id", "", "signing key user id")
|
||||
fs.Int64("max-package-size", 0, "max package size in bytes")
|
||||
fs.Bool("open-registration", true, "allow new account registration")
|
||||
fs.Bool("tls-enabled", false, "enable TLS")
|
||||
fs.String("tls-cert", "", "TLS cert path")
|
||||
fs.String("tls-key", "", "TLS key path")
|
||||
if err := fs.Parse(a.args); err != nil {
|
||||
return config.Config{}, err
|
||||
}
|
||||
|
||||
flags := map[string]string{}
|
||||
fs.Visit(func(f *flag.Flag) {
|
||||
flags[f.Name] = f.Value.String()
|
||||
})
|
||||
return config.Load(*configPath, flags)
|
||||
}
|
||||
|
||||
// ensureDirs creates the store, database, and packages directories.
|
||||
func ensureDirs(cfg config.Config) error {
|
||||
for _, dir := range []string{cfg.StoreDir, filepath.Dir(cfg.DBPath), cfg.PackagesDir} {
|
||||
if dir == "" {
|
||||
continue
|
||||
}
|
||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||
return fmt.Errorf("mkdir %s: %w", dir, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ensureSigningKey loads the default key from the DB or generates one.
|
||||
func ensureSigningKey(ctx context.Context, st *store.Store, cfg config.Config) (*gpg.Key, error) {
|
||||
existing, err := st.GetDefaultGPGKey(ctx)
|
||||
if err == nil {
|
||||
key, err := gpg.ParseArmoredPrivate(existing.PrivateKeyArmored)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parse stored key: %w", err)
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
if !errors.Is(err, store.ErrNotFound) {
|
||||
return nil, fmt.Errorf("load key: %w", err)
|
||||
}
|
||||
|
||||
key, err := gpg.GenerateKey(cfg.SigningKeyUserID, cfg.SigningKeyBits)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("generate key: %w", err)
|
||||
}
|
||||
pub, err := key.ArmoredPublic()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
priv, err := key.ArmoredPrivate()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if _, err := st.SaveGPGKey(ctx, key.Fingerprint, key.UserID, pub, priv, true); err != nil {
|
||||
return nil, fmt.Errorf("save key: %w", err)
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
// shutdownTimeout for graceful HTTP shutdown.
|
||||
const shutdownTimeout = 30 * time.Second
|
||||
@@ -0,0 +1,293 @@
|
||||
// Package aptrepo implements the APT repository endpoint served under /apt/:repo.
|
||||
// It generates Release/InRelease/Packages indices on demand from the database
|
||||
// (cached in memory) and streams .deb files from the content-addressed store.
|
||||
// Public repositories allow anonymous reads; private repositories require HTTP
|
||||
// Basic auth where the password is an API token.
|
||||
package aptrepo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
|
||||
"urapt/server/auth"
|
||||
"urapt/server/cache"
|
||||
"urapt/server/middleware"
|
||||
"urapt/server/store"
|
||||
"urapt/shared/apt"
|
||||
"urapt/shared/gpg"
|
||||
"urapt/shared/httputil"
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// APTRepo is the APT endpoint handler.
|
||||
type APTRepo struct {
|
||||
Store *store.Store
|
||||
Auth *auth.Service
|
||||
Signer *gpg.Key
|
||||
Cache *cache.IndexCache
|
||||
}
|
||||
|
||||
// New returns the APT endpoint http.Handler (to be mounted at /apt).
|
||||
func New(st *store.Store, authSvc *auth.Service, signer *gpg.Key, c *cache.IndexCache) http.Handler {
|
||||
a := &APTRepo{Store: st, Auth: authSvc, Signer: signer, Cache: c}
|
||||
|
||||
r := chi.NewRouter()
|
||||
r.Use(middleware.Recover)
|
||||
r.Use(middleware.Log)
|
||||
|
||||
r.Get("/{repo}/dists/{suite}/InRelease", a.InRelease)
|
||||
r.Get("/{repo}/dists/{suite}/Release", a.Release)
|
||||
r.Get("/{repo}/dists/{suite}/Release.gpg", a.ReleaseGpg)
|
||||
r.Get("/{repo}/dists/{suite}/{component}/{binary}/Packages", a.Packages)
|
||||
r.Get("/{repo}/dists/{suite}/{component}/{binary}/Packages.gz", a.PackagesGz)
|
||||
r.Get("/{repo}/dists/{suite}/{component}/{binary}/Packages.xz", a.PackagesXz)
|
||||
r.Get("/{repo}/pool/{component}/{letter}/{src}/{filename}", a.Pool)
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
// authorize loads the repository and enforces read access. For public repos
|
||||
// anonymous access is allowed; for private repos HTTP Basic (password = token)
|
||||
// is required. Returns the repo and true on success; on failure an error has
|
||||
// been written.
|
||||
func (a *APTRepo) authorize(w http.ResponseWriter, r *http.Request) (*models.Repository, bool) {
|
||||
repo, err := a.Store.GetRepositoryByName(r.Context(), r.PathValue("repo"))
|
||||
if err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "repository not found")
|
||||
return nil, false
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to read repository")
|
||||
return nil, false
|
||||
}
|
||||
|
||||
if repo.Visibility == models.VisibilityPublic {
|
||||
return repo, true
|
||||
}
|
||||
|
||||
// Private repository: require HTTP Basic with token as password.
|
||||
id, err := a.Auth.ResolveBasic(r.Context(), r.Header)
|
||||
if err != nil {
|
||||
httputil.ChallengeBasic(w, "urapt "+repo.Name)
|
||||
return nil, false
|
||||
}
|
||||
ok, err := a.Auth.CanRead(r.Context(), id.User, repo)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "permission check failed")
|
||||
return nil, false
|
||||
}
|
||||
if !ok {
|
||||
httputil.ChallengeBasic(w, "urapt "+repo.Name)
|
||||
return nil, false
|
||||
}
|
||||
return repo, true
|
||||
}
|
||||
|
||||
// loadSuite returns the generated indices for (repo, suite), building and
|
||||
// caching them on first access or after invalidation.
|
||||
func (a *APTRepo) loadSuite(r *http.Request, repo *models.Repository, suiteName string) (*apt.Indices, error) {
|
||||
if idx := a.Cache.Get(repo.ID, suiteName); idx != nil {
|
||||
return idx, nil
|
||||
}
|
||||
|
||||
dist, err := a.Store.GetDistributionByName(r.Context(), repo.ID, suiteName)
|
||||
if err != nil {
|
||||
return nil, errMiss{err}
|
||||
}
|
||||
comps, err := a.Store.ListComponents(r.Context(), dist.ID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
arches, err := a.Store.ListArchitectures(r.Context(), dist.ID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
suitePkgs, err := a.Store.ListSuitePackages(r.Context(), repo.ID, dist.ID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
componentNames := make([]string, 0, len(comps))
|
||||
for _, c := range comps {
|
||||
componentNames = append(componentNames, c.Name)
|
||||
}
|
||||
archNames := make([]string, 0, len(arches))
|
||||
for _, a2 := range arches {
|
||||
archNames = append(archNames, a2.Name)
|
||||
}
|
||||
rows := make([]apt.PackageRow, 0, len(suitePkgs))
|
||||
for _, sp := range suitePkgs {
|
||||
rows = append(rows, apt.PackageRow{
|
||||
Component: sp.ComponentName,
|
||||
Name: sp.Name,
|
||||
Version: sp.Version,
|
||||
Architecture: sp.Architecture,
|
||||
PoolPath: sp.PoolPath,
|
||||
Size: sp.Size,
|
||||
MD5sum: sp.MD5sum,
|
||||
SHA1: sp.SHA1,
|
||||
SHA256: sp.SHA256,
|
||||
DescriptionMD5: sp.DescriptionMD5,
|
||||
RawControl: sp.RawControl,
|
||||
})
|
||||
}
|
||||
|
||||
suite := &apt.Suite{
|
||||
Origin: "urapt " + repo.Name,
|
||||
Label: "urapt " + repo.Name,
|
||||
Suite: dist.Name,
|
||||
Description: repo.Description,
|
||||
Components: componentNames,
|
||||
Architectures: archNames,
|
||||
Packages: rows,
|
||||
}
|
||||
idx, err := apt.Generate(suite, a.Signer)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
a.Cache.Put(repo.ID, suiteName, idx)
|
||||
return idx, nil
|
||||
}
|
||||
|
||||
// errMiss wraps a not-found error from a sub-load.
|
||||
type errMiss struct{ err error }
|
||||
|
||||
func (e errMiss) Error() string { return e.err.Error() }
|
||||
func (e errMiss) Unwrap() error { return e.err }
|
||||
|
||||
// indicesOrWrite loads indices and writes a 404/500 on failure.
|
||||
func (a *APTRepo) indicesOrWrite(w http.ResponseWriter, r *http.Request, repo *models.Repository, suite string) *apt.Indices {
|
||||
idx, err := a.loadSuite(r, repo, suite)
|
||||
if err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "suite not found")
|
||||
return nil
|
||||
}
|
||||
slog.Error("apt generate failed", "err", err, "repo", repo.Name, "suite", suite)
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to generate indices")
|
||||
return nil
|
||||
}
|
||||
return idx
|
||||
}
|
||||
|
||||
// Release serves the unsigned Release file.
|
||||
func (a *APTRepo) Release(w http.ResponseWriter, r *http.Request) {
|
||||
repo, ok := a.authorize(w, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
idx := a.indicesOrWrite(w, r, repo, r.PathValue("suite"))
|
||||
if idx == nil {
|
||||
return
|
||||
}
|
||||
writeBytes(w, "text/plain; charset=utf-8", idx.Release)
|
||||
}
|
||||
|
||||
// InRelease serves the clearsigned InRelease file.
|
||||
func (a *APTRepo) InRelease(w http.ResponseWriter, r *http.Request) {
|
||||
repo, ok := a.authorize(w, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
idx := a.indicesOrWrite(w, r, repo, r.PathValue("suite"))
|
||||
if idx == nil {
|
||||
return
|
||||
}
|
||||
if len(idx.InRelease) == 0 {
|
||||
httputil.WriteError(w, http.StatusServiceUnavailable, httputil.CodeInternal, "indices not signed")
|
||||
return
|
||||
}
|
||||
writeBytes(w, "text/plain; charset=utf-8", idx.InRelease)
|
||||
}
|
||||
|
||||
// ReleaseGpg serves the detached signature Release.gpg.
|
||||
func (a *APTRepo) ReleaseGpg(w http.ResponseWriter, r *http.Request) {
|
||||
repo, ok := a.authorize(w, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
idx := a.indicesOrWrite(w, r, repo, r.PathValue("suite"))
|
||||
if idx == nil {
|
||||
return
|
||||
}
|
||||
if len(idx.ReleaseGpg) == 0 {
|
||||
httputil.WriteError(w, http.StatusServiceUnavailable, httputil.CodeInternal, "indices not signed")
|
||||
return
|
||||
}
|
||||
writeBytes(w, "application/pgp-signature", idx.ReleaseGpg)
|
||||
}
|
||||
|
||||
// Packages serves the Packages index for a (component, arch).
|
||||
func (a *APTRepo) Packages(w http.ResponseWriter, r *http.Request) {
|
||||
a.servePackages(w, r, "")
|
||||
}
|
||||
func (a *APTRepo) PackagesGz(w http.ResponseWriter, r *http.Request) {
|
||||
a.servePackages(w, r, "gz")
|
||||
}
|
||||
func (a *APTRepo) PackagesXz(w http.ResponseWriter, r *http.Request) {
|
||||
a.servePackages(w, r, "xz")
|
||||
}
|
||||
|
||||
func (a *APTRepo) servePackages(w http.ResponseWriter, r *http.Request, encoding string) {
|
||||
repo, ok := a.authorize(w, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
idx := a.indicesOrWrite(w, r, repo, r.PathValue("suite"))
|
||||
if idx == nil {
|
||||
return
|
||||
}
|
||||
binary := r.PathValue("binary")
|
||||
if !strings.HasPrefix(binary, "binary-") {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "not found")
|
||||
return
|
||||
}
|
||||
arch := strings.TrimPrefix(binary, "binary-")
|
||||
relPath := r.PathValue("component") + "/binary-" + arch + "/Packages"
|
||||
var data []byte
|
||||
var contentType string
|
||||
switch encoding {
|
||||
case "":
|
||||
data = idx.Packages[relPath]
|
||||
contentType = "text/plain; charset=utf-8"
|
||||
case "gz":
|
||||
data = idx.PackagesGz[relPath+".gz"]
|
||||
contentType = "application/gzip"
|
||||
case "xz":
|
||||
data = idx.PackagesXz[relPath+".xz"]
|
||||
contentType = "application/x-xz"
|
||||
}
|
||||
if data == nil {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "index not found")
|
||||
return
|
||||
}
|
||||
writeBytes(w, contentType, data)
|
||||
}
|
||||
|
||||
// Pool streams a .deb file from the content-addressed store.
|
||||
func (a *APTRepo) Pool(w http.ResponseWriter, r *http.Request) {
|
||||
repo, ok := a.authorize(w, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
poolPath := strings.Join([]string{
|
||||
"pool", r.PathValue("component"), r.PathValue("letter"), r.PathValue("src"), r.PathValue("filename"),
|
||||
}, "/")
|
||||
pkg, err := a.Store.GetPackageByPoolPath(r.Context(), repo.ID, poolPath)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "package not found")
|
||||
return
|
||||
}
|
||||
http.ServeFile(w, r, pkg.Filename)
|
||||
}
|
||||
|
||||
// writeBytes writes raw bytes with a content type and 200 status.
|
||||
func writeBytes(w http.ResponseWriter, contentType string, data []byte) {
|
||||
w.Header().Set("Content-Type", contentType)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write(data)
|
||||
}
|
||||
@@ -0,0 +1,195 @@
|
||||
package aptrepo_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
|
||||
"urapt/server/aptrepo"
|
||||
"urapt/server/auth"
|
||||
"urapt/server/cache"
|
||||
"urapt/server/restapi"
|
||||
"urapt/server/store"
|
||||
"urapt/shared/config"
|
||||
"urapt/shared/db"
|
||||
"urapt/shared/gpg"
|
||||
)
|
||||
|
||||
// newServer spins up a REST + APT server backed by a temp DB and store, with a
|
||||
// pre-created owner token. It returns the server, the owner token, and the
|
||||
// armored public key.
|
||||
func newServer(t *testing.T) (*httptest.Server, string, string) {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
database, err := db.Open(filepath.Join(dir, "test.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("db open: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { database.Close() })
|
||||
st := store.New(database)
|
||||
authSvc := auth.NewService(st)
|
||||
|
||||
key, err := gpg.GenerateKey("urapt-test <test>", 2048)
|
||||
if err != nil {
|
||||
t.Fatalf("gpg key: %v", err)
|
||||
}
|
||||
sp := &testSigner{key: key}
|
||||
cfg := config.Defaults
|
||||
cfg.StoreDir = dir
|
||||
cfg.PackagesDir = filepath.Join(dir, "packages")
|
||||
cfg.DBPath = filepath.Join(dir, "test.db")
|
||||
if err := mkdirAll(cfg.PackagesDir); err != nil {
|
||||
t.Fatalf("mkdir pkg dir: %v", err)
|
||||
}
|
||||
|
||||
idxCache := cache.New()
|
||||
root := chi.NewRouter()
|
||||
root.Mount("/api/v1", restapi.New(st, authSvc, sp, &cfg, idxCache))
|
||||
root.Mount("/apt", aptrepo.New(st, authSvc, key, idxCache))
|
||||
srv := httptest.NewServer(root)
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
// Register owner.
|
||||
body := post(t, srv, "/api/v1/auth/register", "", map[string]string{"username": "owner", "password": "supersecret"})
|
||||
token := str(body, "token")
|
||||
pub, err := key.ArmoredPublic()
|
||||
if err != nil {
|
||||
t.Fatalf("armored public: %v", err)
|
||||
}
|
||||
return srv, token, pub
|
||||
}
|
||||
|
||||
type testSigner struct{ key *gpg.Key }
|
||||
|
||||
func (s *testSigner) PublicKeyArmored() (string, error) { return s.key.ArmoredPublic() }
|
||||
func (s *testSigner) Fingerprint() string { return s.key.Fingerprint }
|
||||
|
||||
func post(t *testing.T, srv *httptest.Server, path, token string, body any) []byte {
|
||||
t.Helper()
|
||||
b := jsonMarshal(body)
|
||||
req, _ := http.NewRequest("POST", srv.URL+path, bytes.NewReader(b))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
}
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatalf("post %s: %v", path, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
return readAll(resp.Body)
|
||||
}
|
||||
|
||||
func get(t *testing.T, srv *httptest.Server, path, token string) (int, []byte) {
|
||||
t.Helper()
|
||||
req, _ := http.NewRequest("GET", srv.URL+path, nil)
|
||||
if token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
}
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatalf("get %s: %v", path, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
return resp.StatusCode, readAll(resp.Body)
|
||||
}
|
||||
|
||||
func setupRepo(t *testing.T, srv *httptest.Server, token string) {
|
||||
post(t, srv, "/api/v1/repositories", token, map[string]any{"name": "myrepo", "visibility": "public"})
|
||||
post(t, srv, "/api/v1/repositories/myrepo/distributions", token, map[string]string{"name": "stable"})
|
||||
post(t, srv, "/api/v1/repositories/myrepo/distributions/stable/components", token, map[string]string{"name": "main"})
|
||||
post(t, srv, "/api/v1/repositories/myrepo/distributions/stable/architectures", token, map[string]string{"name": "amd64"})
|
||||
}
|
||||
|
||||
func TestAPTIndicesAndPool(t *testing.T) {
|
||||
srv, token, pub := newServer(t)
|
||||
setupRepo(t, srv, token)
|
||||
|
||||
// Push a real .deb via the REST API.
|
||||
debBytes := buildDeb("foo", "1.0", "amd64")
|
||||
upload(t, srv, token, "myrepo", "stable", "main", "foo_1.0_amd64.deb", debBytes)
|
||||
|
||||
// Fetch the Packages index.
|
||||
code, body := get(t, srv, "/apt/myrepo/dists/stable/main/binary-amd64/Packages", "")
|
||||
if code != 200 {
|
||||
t.Fatalf("Packages status %d body %s", code, body)
|
||||
}
|
||||
if !bytes.Contains(body, []byte("Package: foo")) || !bytes.Contains(body, []byte("Filename: pool/main/f/foo/foo_1.0_amd64.deb")) {
|
||||
t.Fatalf("Packages index missing entry:\n%s", body)
|
||||
}
|
||||
|
||||
// Fetch Release.
|
||||
code, rel := get(t, srv, "/apt/myrepo/dists/stable/Release", "")
|
||||
if code != 200 {
|
||||
t.Fatalf("Release status %d", code)
|
||||
}
|
||||
if !bytes.Contains(rel, []byte("Suite: stable")) || !bytes.Contains(rel, []byte("Components: main")) {
|
||||
t.Fatalf("Release missing fields:\n%s", rel)
|
||||
}
|
||||
|
||||
// Fetch InRelease (clearsigned) and verify against the public key.
|
||||
code, inrel := get(t, srv, "/apt/myrepo/dists/stable/InRelease", "")
|
||||
if code != 200 {
|
||||
t.Fatalf("InRelease status %d", code)
|
||||
}
|
||||
if !bytes.Contains(inrel, []byte("BEGIN PGP SIGNED MESSAGE")) {
|
||||
t.Fatalf("InRelease not clearsigned:\n%s", inrel)
|
||||
}
|
||||
if _, err := gpg.VerifyClearSign(pub, inrel); err != nil {
|
||||
t.Fatalf("verify InRelease: %v", err)
|
||||
}
|
||||
|
||||
// Fetch Release.gpg and verify as a detached signature of Release.
|
||||
code, sig := get(t, srv, "/apt/myrepo/dists/stable/Release.gpg", "")
|
||||
if code != 200 {
|
||||
t.Fatalf("Release.gpg status %d", code)
|
||||
}
|
||||
if err := gpg.VerifyDetached(pub, rel, sig); err != nil {
|
||||
t.Fatalf("verify Release.gpg: %v", err)
|
||||
}
|
||||
|
||||
// Fetch the pool .deb and compare bytes.
|
||||
code, deb := get(t, srv, "/apt/myrepo/pool/main/f/foo/foo_1.0_amd64.deb", "")
|
||||
if code != 200 {
|
||||
t.Fatalf("pool status %d", code)
|
||||
}
|
||||
if !bytes.Equal(deb, debBytes) {
|
||||
t.Fatalf("pool bytes mismatch (%d vs %d)", len(deb), len(debBytes))
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrivateRepoRequiresAuth(t *testing.T) {
|
||||
srv, token, _ := newServer(t)
|
||||
// private repo
|
||||
post(t, srv, "/api/v1/repositories", token, map[string]any{"name": "priv", "visibility": "private"})
|
||||
post(t, srv, "/api/v1/repositories/priv/distributions", token, map[string]string{"name": "stable"})
|
||||
post(t, srv, "/api/v1/repositories/priv/distributions/stable/components", token, map[string]string{"name": "main"})
|
||||
post(t, srv, "/api/v1/repositories/priv/distributions/stable/architectures", token, map[string]string{"name": "amd64"})
|
||||
upload(t, srv, token, "priv", "stable", "main", "foo_1.0_amd64.deb", buildDeb("foo", "1.0", "amd64"))
|
||||
|
||||
// anonymous → 401
|
||||
code, _ := get(t, srv, "/apt/priv/dists/stable/Release", "")
|
||||
if code != 401 {
|
||||
t.Fatalf("anonymous private repo should be 401, got %d", code)
|
||||
}
|
||||
|
||||
// with token as basic auth password → 200
|
||||
req, _ := http.NewRequest("GET", srv.URL+"/apt/priv/dists/stable/Release", nil)
|
||||
req.SetBasicAuth("owner", token)
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatalf("get: %v", err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode != 200 {
|
||||
t.Fatalf("authenticated private repo should be 200, got %d", resp.StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
// keep context import used
|
||||
var _ = context.Background
|
||||
@@ -0,0 +1,126 @@
|
||||
package aptrepo_test
|
||||
|
||||
import (
|
||||
"archive/tar"
|
||||
"bytes"
|
||||
"compress/gzip"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func jsonMarshal(v any) []byte {
|
||||
b, _ := json.Marshal(v)
|
||||
return b
|
||||
}
|
||||
|
||||
func readAll(r io.Reader) []byte {
|
||||
b, _ := io.ReadAll(r)
|
||||
return b
|
||||
}
|
||||
|
||||
func mkdirAll(dir string) error { return os.MkdirAll(dir, 0o755) }
|
||||
|
||||
func str(body []byte, key string) string {
|
||||
var m map[string]any
|
||||
_ = json.Unmarshal(body, &m)
|
||||
if v, ok := m[key].(string); ok {
|
||||
return v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// buildDeb constructs a minimal valid .deb with the given package/version/arch.
|
||||
func buildDeb(name, version, arch string) []byte {
|
||||
control := "Package: " + name + "\nVersion: " + version + "\nArchitecture: " + arch + "\nMaintainer: t <t@e>\nDescription: short\n extended\n"
|
||||
var ctrlBuf bytes.Buffer
|
||||
gz := gzip.NewWriter(&ctrlBuf)
|
||||
tw := tar.NewWriter(gz)
|
||||
writeTw(tw, "control", control)
|
||||
tw.Close()
|
||||
gz.Close()
|
||||
|
||||
var dataBuf bytes.Buffer
|
||||
gz2 := gzip.NewWriter(&dataBuf)
|
||||
tw2 := tar.NewWriter(gz2)
|
||||
writeTw(tw2, "usr/share/"+name, "x")
|
||||
tw2.Close()
|
||||
gz2.Close()
|
||||
|
||||
var out bytes.Buffer
|
||||
out.WriteString("!<arch>\n")
|
||||
writeAr(&out, "debian-binary", []byte("2.0\n"))
|
||||
writeAr(&out, "control.tar.gz", ctrlBuf.Bytes())
|
||||
writeAr(&out, "data.tar.gz", dataBuf.Bytes())
|
||||
return out.Bytes()
|
||||
}
|
||||
|
||||
func writeTw(tw *tar.Writer, name, body string) {
|
||||
_ = tw.WriteHeader(&tar.Header{Name: name, Mode: 0o644, Size: int64(len(body)), Typeflag: tar.TypeReg})
|
||||
_, _ = tw.Write([]byte(body))
|
||||
}
|
||||
|
||||
func writeAr(buf *bytes.Buffer, name string, data []byte) {
|
||||
header := make([]byte, 60)
|
||||
for i := range header {
|
||||
header[i] = ' '
|
||||
}
|
||||
copy(header[0:], name+"/")
|
||||
ds := []byte(padNum(len(data), 10))
|
||||
copy(header[48:], ds)
|
||||
header[58] = '`'
|
||||
header[59] = '\n'
|
||||
buf.Write(header)
|
||||
buf.Write(data)
|
||||
if len(data)%2 == 1 {
|
||||
buf.WriteByte('\n')
|
||||
}
|
||||
}
|
||||
|
||||
func padNum(n, width int) string {
|
||||
s := make([]byte, width)
|
||||
for i := range s {
|
||||
s[i] = ' '
|
||||
}
|
||||
digits := []byte(itoa(n))
|
||||
copy(s[len(s)-len(digits):], digits)
|
||||
return string(s)
|
||||
}
|
||||
|
||||
func itoa(n int) string {
|
||||
if n == 0 {
|
||||
return "0"
|
||||
}
|
||||
var b []byte
|
||||
for n > 0 {
|
||||
b = append([]byte{byte('0' + n%10)}, b...)
|
||||
n /= 10
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
|
||||
func upload(t *testing.T, srv *httptest.Server, token, repo, dist, component, filename string, deb []byte) {
|
||||
t.Helper()
|
||||
var buf bytes.Buffer
|
||||
mw := multipart.NewWriter(&buf)
|
||||
_ = mw.WriteField("component", component)
|
||||
fw, _ := mw.CreateFormFile("file", filename)
|
||||
_, _ = fw.Write(deb)
|
||||
_ = mw.Close()
|
||||
|
||||
req, _ := http.NewRequest("POST", srv.URL+"/api/v1/repositories/"+repo+"/distributions/"+dist+"/packages", &buf)
|
||||
req.Header.Set("Content-Type", mw.FormDataContentType())
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatalf("upload: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != 201 {
|
||||
t.Fatalf("upload status %d", resp.StatusCode)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,135 @@
|
||||
// Package auth implements identity resolution (bearer/basic) and the
|
||||
// repository-scoped permission checks used by the REST API and the APT
|
||||
// endpoint.
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"urapt/server/store"
|
||||
"urapt/shared/crypto"
|
||||
"urapt/shared/httputil"
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// Identity is the resolved caller: a user and (for REST) the token used.
|
||||
type Identity struct {
|
||||
User *models.User
|
||||
TokenID string
|
||||
}
|
||||
|
||||
// Service resolves identities and answers permission questions.
|
||||
type Service struct {
|
||||
store *store.Store
|
||||
}
|
||||
|
||||
// NewService constructs an auth Service.
|
||||
func NewService(s *store.Store) *Service {
|
||||
return &Service{store: s}
|
||||
}
|
||||
|
||||
// ErrUnauthenticated is returned when no valid identity can be established.
|
||||
var ErrUnauthenticated = errors.New("unauthenticated")
|
||||
|
||||
// ResolveBearer resolves an "Authorization: Bearer <token>" header to an
|
||||
// identity. Returns ErrUnauthenticated if absent or invalid.
|
||||
func (s *Service) ResolveBearer(ctx context.Context, h http.Header) (*Identity, error) {
|
||||
token, ok := httputil.ParseBearer(h)
|
||||
if !ok {
|
||||
return nil, ErrUnauthenticated
|
||||
}
|
||||
return s.resolveToken(ctx, token)
|
||||
}
|
||||
|
||||
// ResolveBasic resolves an "Authorization: Basic" header where the password is
|
||||
// an API token. Returns ErrUnauthenticated if absent/invalid.
|
||||
func (s *Service) ResolveBasic(ctx context.Context, h http.Header) (*Identity, error) {
|
||||
_, password, ok := httputil.ParseBasic(h)
|
||||
if !ok {
|
||||
return nil, ErrUnauthenticated
|
||||
}
|
||||
return s.resolveToken(ctx, password)
|
||||
}
|
||||
|
||||
func (s *Service) resolveToken(ctx context.Context, token string) (*Identity, error) {
|
||||
if token == "" {
|
||||
return nil, ErrUnauthenticated
|
||||
}
|
||||
hash := crypto.HashToken(token)
|
||||
t, err := s.store.GetTokenByHash(ctx, hash)
|
||||
if err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
return nil, ErrUnauthenticated
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
user, err := s.store.GetUserByID(ctx, t.UserID)
|
||||
if err != nil {
|
||||
return nil, ErrUnauthenticated
|
||||
}
|
||||
_ = s.store.TouchToken(ctx, t.ID)
|
||||
return &Identity{User: user, TokenID: t.ID}, nil
|
||||
}
|
||||
|
||||
// CanRead reports whether user may read (download/list within) repo.
|
||||
func (s *Service) CanRead(ctx context.Context, user *models.User, repo *models.Repository) (bool, error) {
|
||||
if user == nil {
|
||||
return repo.Visibility == models.VisibilityPublic, nil
|
||||
}
|
||||
if user.IsAdmin || repo.OwnerUserID == user.ID {
|
||||
return true, nil
|
||||
}
|
||||
access, ok, err := s.store.GetMemberAccess(ctx, repo.ID, user.ID)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if ok {
|
||||
switch access {
|
||||
case models.AccessRead, models.AccessReadWrite, models.AccessAdmin:
|
||||
return true, nil
|
||||
}
|
||||
}
|
||||
return repo.Visibility == models.VisibilityPublic, nil
|
||||
}
|
||||
|
||||
// CanWrite reports whether user may push packages to repo.
|
||||
func (s *Service) CanWrite(ctx context.Context, user *models.User, repo *models.Repository) (bool, error) {
|
||||
if user == nil {
|
||||
return false, nil
|
||||
}
|
||||
if user.IsAdmin || repo.OwnerUserID == user.ID {
|
||||
return true, nil
|
||||
}
|
||||
access, ok, err := s.store.GetMemberAccess(ctx, repo.ID, user.ID)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if !ok {
|
||||
return false, nil
|
||||
}
|
||||
switch access {
|
||||
case models.AccessWrite, models.AccessReadWrite, models.AccessAdmin:
|
||||
return true, nil
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// CanManage reports whether user may manage repo settings and members.
|
||||
func (s *Service) CanManage(ctx context.Context, user *models.User, repo *models.Repository) (bool, error) {
|
||||
if user == nil {
|
||||
return false, nil
|
||||
}
|
||||
if user.IsAdmin || repo.OwnerUserID == user.ID {
|
||||
return true, nil
|
||||
}
|
||||
access, ok, err := s.store.GetMemberAccess(ctx, repo.ID, user.ID)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if !ok {
|
||||
return false, nil
|
||||
}
|
||||
return access == models.AccessAdmin, nil
|
||||
}
|
||||
@@ -0,0 +1,333 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"urapt/server/store"
|
||||
"urapt/shared/crypto"
|
||||
"urapt/shared/db"
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// newAuthSvc opens a fresh DB and returns an auth Service plus a handle to the
|
||||
// underlying store for seeding users/repos/tokens.
|
||||
func newAuthSvc(t *testing.T) (*Service, *store.Store) {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
database, err := db.Open(filepath.Join(dir, "test.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("db open: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { database.Close() })
|
||||
st := store.New(database)
|
||||
return NewService(st), st
|
||||
}
|
||||
|
||||
// seedUser creates a user via the store and returns it.
|
||||
func seedUser(t *testing.T, st *store.Store, ctx context.Context, username string) *models.User {
|
||||
t.Helper()
|
||||
hash, err := crypto.HashPassword("pw")
|
||||
if err != nil {
|
||||
t.Fatalf("hash: %v", err)
|
||||
}
|
||||
u, _, err := st.CreateUser(ctx, username, hash)
|
||||
if err != nil {
|
||||
t.Fatalf("create user %q: %v", username, err)
|
||||
}
|
||||
return u
|
||||
}
|
||||
|
||||
// seedToken mints a real token for userID and returns the plaintext.
|
||||
func seedToken(t *testing.T, st *store.Store, ctx context.Context, userID, name string) string {
|
||||
t.Helper()
|
||||
tok, hash, prefix, err := crypto.GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("generate: %v", err)
|
||||
}
|
||||
if _, err := st.CreateToken(ctx, userID, name, prefix, hash); err != nil {
|
||||
t.Fatalf("create token: %v", err)
|
||||
}
|
||||
return tok
|
||||
}
|
||||
|
||||
func seedRepo(t *testing.T, st *store.Store, ctx context.Context, name, ownerID string, vis models.Visibility) *models.Repository {
|
||||
t.Helper()
|
||||
r, err := st.CreateRepository(ctx, name, ownerID, vis, "")
|
||||
if err != nil {
|
||||
t.Fatalf("create repo %q: %v", name, err)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// seedAdmin creates a non-bootstrap user and promotes them to server admin,
|
||||
// returning a refreshed user record with IsAdmin=true.
|
||||
func seedAdmin(t *testing.T, st *store.Store, ctx context.Context, username string) *models.User {
|
||||
t.Helper()
|
||||
u := seedUser(t, st, ctx, username)
|
||||
if err := st.UpdateUser(ctx, u.ID, boolPtr(true)); err != nil {
|
||||
t.Fatalf("promote %q: %v", username, err)
|
||||
}
|
||||
refreshed, err := st.GetUserByID(ctx, u.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("refetch admin: %v", err)
|
||||
}
|
||||
return refreshed
|
||||
}
|
||||
|
||||
func bearerHeader(token string) http.Header {
|
||||
h := http.Header{}
|
||||
h.Set("Authorization", "Bearer "+token)
|
||||
return h
|
||||
}
|
||||
|
||||
func basicHeader(username, password string) http.Header {
|
||||
h := http.Header{}
|
||||
enc := base64.StdEncoding.EncodeToString([]byte(username + ":" + password))
|
||||
h.Set("Authorization", "Basic "+enc)
|
||||
return h
|
||||
}
|
||||
|
||||
// --- ResolveBearer ---
|
||||
|
||||
func TestResolveBearer_Valid(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc, st := newAuthSvc(t)
|
||||
u := seedUser(t, st, ctx, "alice")
|
||||
tok := seedToken(t, st, ctx, u.ID, "laptop")
|
||||
|
||||
id, err := svc.ResolveBearer(ctx, bearerHeader(tok))
|
||||
if err != nil {
|
||||
t.Fatalf("resolve: %v", err)
|
||||
}
|
||||
if id.User == nil || id.User.ID != u.ID {
|
||||
t.Fatalf("identity = %+v", id)
|
||||
}
|
||||
if id.TokenID == "" {
|
||||
t.Fatal("token id should be set")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveBearer_MissingHeader(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc, _ := newAuthSvc(t)
|
||||
_, err := svc.ResolveBearer(ctx, http.Header{})
|
||||
if !errors.Is(err, ErrUnauthenticated) {
|
||||
t.Fatalf("expected ErrUnauthenticated, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveBearer_MalformedHeader(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc, _ := newAuthSvc(t)
|
||||
h := http.Header{}
|
||||
h.Set("Authorization", "Bearer") // no token
|
||||
if _, err := svc.ResolveBearer(ctx, h); !errors.Is(err, ErrUnauthenticated) {
|
||||
t.Fatalf("expected ErrUnauthenticated, got %v", err)
|
||||
}
|
||||
h.Set("Authorization", "Basic abc")
|
||||
if _, err := svc.ResolveBearer(ctx, h); !errors.Is(err, ErrUnauthenticated) {
|
||||
t.Fatalf("expected ErrUnauthenticated for non-Bearer scheme, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveBearer_InvalidToken(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc, _ := newAuthSvc(t)
|
||||
if _, err := svc.ResolveBearer(ctx, bearerHeader("urapt_notreal")); !errors.Is(err, ErrUnauthenticated) {
|
||||
t.Fatalf("expected ErrUnauthenticated for bogus token, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveBearer_RevokedToken(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc, st := newAuthSvc(t)
|
||||
u := seedUser(t, st, ctx, "alice")
|
||||
tok := seedToken(t, st, ctx, u.ID, "laptop")
|
||||
// Revoke by hash (simulating logout).
|
||||
_ = st.RevokeTokenByHash(ctx, crypto.HashToken(tok))
|
||||
|
||||
if _, err := svc.ResolveBearer(ctx, bearerHeader(tok)); !errors.Is(err, ErrUnauthenticated) {
|
||||
t.Fatalf("expected ErrUnauthenticated for revoked token, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// --- ResolveBasic ---
|
||||
|
||||
func TestResolveBasic_Valid(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc, st := newAuthSvc(t)
|
||||
u := seedUser(t, st, ctx, "alice")
|
||||
tok := seedToken(t, st, ctx, u.ID, "laptop")
|
||||
|
||||
// Username is ignored; password is the token.
|
||||
id, err := svc.ResolveBasic(ctx, basicHeader("anything", tok))
|
||||
if err != nil {
|
||||
t.Fatalf("resolve: %v", err)
|
||||
}
|
||||
if id.User.ID != u.ID {
|
||||
t.Fatalf("user id mismatch")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveBasic_MissingAndMalformed(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc, _ := newAuthSvc(t)
|
||||
if _, err := svc.ResolveBasic(ctx, http.Header{}); !errors.Is(err, ErrUnauthenticated) {
|
||||
t.Fatalf("missing header: expected ErrUnauthenticated, got %v", err)
|
||||
}
|
||||
h := http.Header{}
|
||||
h.Set("Authorization", "Basic not-base64!!!")
|
||||
if _, err := svc.ResolveBasic(ctx, h); !errors.Is(err, ErrUnauthenticated) {
|
||||
t.Fatalf("bad base64: expected ErrUnauthenticated, got %v", err)
|
||||
}
|
||||
h.Set("Authorization", "Basic "+base64.StdEncoding.EncodeToString([]byte("noseparator")))
|
||||
if _, err := svc.ResolveBasic(ctx, h); !errors.Is(err, ErrUnauthenticated) {
|
||||
t.Fatalf("no colon: expected ErrUnauthenticated, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// --- CanRead ---
|
||||
|
||||
func TestCanRead(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc, st := newAuthSvc(t)
|
||||
owner := seedUser(t, st, ctx, "owner")
|
||||
member := seedUser(t, st, ctx, "member")
|
||||
nonMember := seedUser(t, st, ctx, "stranger")
|
||||
admin := seedAdmin(t, st, ctx, "admin")
|
||||
|
||||
pubRepo := seedRepo(t, st, ctx, "pub", owner.ID, models.VisibilityPublic)
|
||||
privRepo := seedRepo(t, st, ctx, "priv", owner.ID, models.VisibilityPrivate)
|
||||
_ = st.AddMember(ctx, privRepo.ID, member.ID, models.AccessRead)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
user *models.User
|
||||
repo *models.Repository
|
||||
want bool
|
||||
}{
|
||||
{"nil user, public repo", nil, pubRepo, true},
|
||||
{"nil user, private repo", nil, privRepo, false},
|
||||
{"admin on private", admin, privRepo, true},
|
||||
{"owner on private", owner, privRepo, true},
|
||||
{"member(read) on private", member, privRepo, true},
|
||||
{"non-member on public", nonMember, pubRepo, true},
|
||||
{"non-member on private", nonMember, privRepo, false},
|
||||
}
|
||||
for _, c := range cases {
|
||||
got, err := svc.CanRead(ctx, c.user, c.repo)
|
||||
if err != nil {
|
||||
t.Fatalf("%s: error %v", c.name, err)
|
||||
}
|
||||
if got != c.want {
|
||||
t.Errorf("%s: got %v want %v", c.name, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCanRead_WriteOnlyMemberCanRead(t *testing.T) {
|
||||
// A write-only member should still be able to read (the AccessWrite grant
|
||||
// does not include read in the switch in CanRead, so this confirms the
|
||||
// public-fallback behavior: a private repo would deny a write-only member
|
||||
// read access).
|
||||
ctx := context.Background()
|
||||
svc, st := newAuthSvc(t)
|
||||
owner := seedUser(t, st, ctx, "owner")
|
||||
writer := seedUser(t, st, ctx, "writer")
|
||||
privRepo := seedRepo(t, st, ctx, "priv", owner.ID, models.VisibilityPrivate)
|
||||
_ = st.AddMember(ctx, privRepo.ID, writer.ID, models.AccessWrite)
|
||||
|
||||
got, _ := svc.CanRead(ctx, writer, privRepo)
|
||||
if got {
|
||||
t.Fatal("write-only member should NOT be able to read a private repo")
|
||||
}
|
||||
}
|
||||
|
||||
// --- CanWrite ---
|
||||
|
||||
func TestCanWrite(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc, st := newAuthSvc(t)
|
||||
owner := seedUser(t, st, ctx, "owner")
|
||||
reader := seedUser(t, st, ctx, "reader")
|
||||
writer := seedUser(t, st, ctx, "writer")
|
||||
rw := seedUser(t, st, ctx, "rw")
|
||||
repoAdmin := seedUser(t, st, ctx, "repoadmin")
|
||||
nonMember := seedUser(t, st, ctx, "stranger")
|
||||
admin := seedAdmin(t, st, ctx, "admin")
|
||||
|
||||
repo := seedRepo(t, st, ctx, "repo", owner.ID, models.VisibilityPublic)
|
||||
_ = st.AddMember(ctx, repo.ID, reader.ID, models.AccessRead)
|
||||
_ = st.AddMember(ctx, repo.ID, writer.ID, models.AccessWrite)
|
||||
_ = st.AddMember(ctx, repo.ID, rw.ID, models.AccessReadWrite)
|
||||
_ = st.AddMember(ctx, repo.ID, repoAdmin.ID, models.AccessAdmin)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
user *models.User
|
||||
want bool
|
||||
}{
|
||||
{"nil user", nil, false},
|
||||
{"admin", admin, true},
|
||||
{"owner", owner, true},
|
||||
{"read member", reader, false},
|
||||
{"write member", writer, true},
|
||||
{"read-write member", rw, true},
|
||||
{"repo admin member", repoAdmin, true},
|
||||
{"non-member", nonMember, false},
|
||||
}
|
||||
for _, c := range cases {
|
||||
got, err := svc.CanWrite(ctx, c.user, repo)
|
||||
if err != nil {
|
||||
t.Fatalf("%s: error %v", c.name, err)
|
||||
}
|
||||
if got != c.want {
|
||||
t.Errorf("%s: got %v want %v", c.name, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// --- CanManage ---
|
||||
|
||||
func TestCanManage(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc, st := newAuthSvc(t)
|
||||
owner := seedUser(t, st, ctx, "owner")
|
||||
reader := seedUser(t, st, ctx, "reader")
|
||||
repoAdmin := seedUser(t, st, ctx, "repoadmin")
|
||||
nonMember := seedUser(t, st, ctx, "stranger")
|
||||
admin := seedAdmin(t, st, ctx, "admin")
|
||||
|
||||
repo := seedRepo(t, st, ctx, "repo", owner.ID, models.VisibilityPublic)
|
||||
_ = st.AddMember(ctx, repo.ID, reader.ID, models.AccessRead)
|
||||
_ = st.AddMember(ctx, repo.ID, repoAdmin.ID, models.AccessAdmin)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
user *models.User
|
||||
want bool
|
||||
}{
|
||||
{"nil user", nil, false},
|
||||
{"admin", admin, true},
|
||||
{"owner", owner, true},
|
||||
{"read member", reader, false},
|
||||
{"repo admin member", repoAdmin, true},
|
||||
{"non-member", nonMember, false},
|
||||
}
|
||||
for _, c := range cases {
|
||||
got, err := svc.CanManage(ctx, c.user, repo)
|
||||
if err != nil {
|
||||
t.Fatalf("%s: error %v", c.name, err)
|
||||
}
|
||||
if got != c.want {
|
||||
t.Errorf("%s: got %v want %v", c.name, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func boolPtr(b bool) *bool { return &b }
|
||||
Vendored
+79
@@ -0,0 +1,79 @@
|
||||
// Package cache defines the in-memory APT index cache contract used by the
|
||||
// REST API (to invalidate on mutation) and the APT endpoint (to serve cached
|
||||
// indices). The implementation lives in this package; it is regenerated lazily
|
||||
// after invalidation or restart.
|
||||
package cache
|
||||
|
||||
import (
|
||||
"sync"
|
||||
|
||||
"urapt/shared/apt"
|
||||
)
|
||||
|
||||
// IndexCache holds generated APT indices per (repository, suite), keyed and
|
||||
// versioned so that mutations invalidate lazily.
|
||||
type IndexCache struct {
|
||||
mu sync.Mutex
|
||||
entry map[string]*entry
|
||||
}
|
||||
|
||||
type entry struct {
|
||||
indices *apt.Indices
|
||||
suite *apt.Suite
|
||||
gen int64
|
||||
dirty bool
|
||||
}
|
||||
|
||||
// New constructs an empty IndexCache.
|
||||
func New() *IndexCache {
|
||||
return &IndexCache{entry: map[string]*entry{}}
|
||||
}
|
||||
|
||||
// key builds the cache key.
|
||||
func key(repoID, suite string) string { return repoID + "/" + suite }
|
||||
|
||||
// Get returns the cached indices for (repoID, suite), or nil if not present
|
||||
// or marked dirty. The suite is returned so the caller can rebuild it.
|
||||
func (c *IndexCache) Get(repoID, suite string) *apt.Indices {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
e := c.entry[key(repoID, suite)]
|
||||
if e == nil || e.dirty {
|
||||
return nil
|
||||
}
|
||||
return e.indices
|
||||
}
|
||||
|
||||
// Put stores freshly generated indices for (repoID, suite).
|
||||
func (c *IndexCache) Put(repoID, suite string, idx *apt.Indices) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
e := c.entry[key(repoID, suite)]
|
||||
if e == nil {
|
||||
e = &entry{}
|
||||
c.entry[key(repoID, suite)] = e
|
||||
}
|
||||
e.indices = idx
|
||||
e.dirty = false
|
||||
}
|
||||
|
||||
// Invalidate marks one (repository, suite) as stale; the next read rebuilds.
|
||||
func (c *IndexCache) Invalidate(repoID, suite string) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if e := c.entry[key(repoID, suite)]; e != nil {
|
||||
e.dirty = true
|
||||
}
|
||||
}
|
||||
|
||||
// InvalidateRepo marks every suite under a repository as stale.
|
||||
func (c *IndexCache) InvalidateRepo(repoID string) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
prefix := repoID + "/"
|
||||
for k, e := range c.entry {
|
||||
if len(k) > len(prefix) && k[:len(prefix)] == prefix {
|
||||
e.dirty = true
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,107 @@
|
||||
// Package middleware provides HTTP middleware shared by the REST API and the
|
||||
// APT endpoint: panic recovery, request logging, and bearer/basic identity
|
||||
// injection.
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"runtime/debug"
|
||||
|
||||
"urapt/server/auth"
|
||||
"urapt/shared/httputil"
|
||||
)
|
||||
|
||||
// ctxKey is an unexported key type for context values.
|
||||
type ctxKey int
|
||||
|
||||
const (
|
||||
keyIdentity ctxKey = iota
|
||||
)
|
||||
|
||||
// IdentityFromContext returns the identity previously attached by RequireBearer
|
||||
// or OptionalBearer, or nil.
|
||||
func IdentityFromContext(ctx context.Context) *auth.Identity {
|
||||
v, _ := ctx.Value(keyIdentity).(*auth.Identity)
|
||||
return v
|
||||
}
|
||||
|
||||
// withIdentity stores id in the context.
|
||||
func withIdentity(ctx context.Context, id *auth.Identity) context.Context {
|
||||
return context.WithValue(ctx, keyIdentity, id)
|
||||
}
|
||||
|
||||
// Recover catches panics and renders a uniform 500.
|
||||
func Recover(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
slog.Error("panic", "err", rec, "stack", string(debug.Stack()))
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "internal server error")
|
||||
}
|
||||
}()
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
// Log logs each request using slog.
|
||||
func Log(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
rw := &statusRecorder{ResponseWriter: w, status: http.StatusOK}
|
||||
next.ServeHTTP(rw, r)
|
||||
slog.Info("http", "method", r.Method, "path", r.URL.Path, "status", rw.status)
|
||||
})
|
||||
}
|
||||
|
||||
type statusRecorder struct {
|
||||
http.ResponseWriter
|
||||
status int
|
||||
}
|
||||
|
||||
func (s *statusRecorder) WriteHeader(code int) {
|
||||
s.status = code
|
||||
s.ResponseWriter.WriteHeader(code)
|
||||
}
|
||||
|
||||
// RequireBearer resolves a bearer token; on failure it renders 401. On success
|
||||
// the identity is attached to the request context.
|
||||
func RequireBearer(svc *auth.Service) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := svc.ResolveBearer(r.Context(), r.Header)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusUnauthorized, httputil.CodeUnauthorized, "authentication required")
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r.WithContext(withIdentity(r.Context(), id)))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// OptionalBearer resolves a bearer token if present but never blocks; the
|
||||
// identity (possibly nil) is attached to the context.
|
||||
func OptionalBearer(svc *auth.Service) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
id, _ := svc.ResolveBearer(r.Context(), r.Header)
|
||||
next.ServeHTTP(w, r.WithContext(withIdentity(r.Context(), id)))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// RequireBasic resolves HTTP Basic auth (password = API token); on failure it
|
||||
// issues a 401 with a WWW-Authenticate challenge. Used by the APT endpoint for
|
||||
// private repositories.
|
||||
func RequireBasic(svc *auth.Service, realm string) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := svc.ResolveBasic(r.Context(), r.Header)
|
||||
if err != nil {
|
||||
httputil.ChallengeBasic(w, realm)
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r.WithContext(withIdentity(r.Context(), id)))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
// Package restapi implements the urapt REST API: handlers, routing, and
|
||||
// request validation. It is mounted under /api/v1 on the server.
|
||||
package restapi
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"regexp"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
|
||||
"urapt/server/auth"
|
||||
"urapt/server/cache"
|
||||
"urapt/server/middleware"
|
||||
"urapt/server/store"
|
||||
"urapt/shared/config"
|
||||
"urapt/shared/httputil"
|
||||
)
|
||||
|
||||
// SignerProvider returns the server's current signing key. It is supplied by
|
||||
// the app; the APT endpoint also uses it to sign Release files.
|
||||
type SignerProvider interface {
|
||||
PublicKeyArmored() (string, error)
|
||||
Fingerprint() string
|
||||
}
|
||||
|
||||
// API holds the dependencies shared by all REST handlers.
|
||||
type API struct {
|
||||
Store *store.Store
|
||||
Auth *auth.Service
|
||||
Signer SignerProvider
|
||||
Config *config.Config
|
||||
Cache *cache.IndexCache
|
||||
}
|
||||
|
||||
// New constructs the API and returns its http.Handler (the /api/v1 router).
|
||||
func New(st *store.Store, authSvc *auth.Service, signer SignerProvider, cfg *config.Config, c *cache.IndexCache) http.Handler {
|
||||
api := &API{Store: st, Auth: authSvc, Signer: signer, Config: cfg, Cache: c}
|
||||
|
||||
r := chi.NewRouter()
|
||||
r.Use(middleware.Recover)
|
||||
r.Use(middleware.Log)
|
||||
|
||||
r.Get("/server/info", api.ServerInfo)
|
||||
r.Get("/server/pubkey", api.ServerPubkey)
|
||||
|
||||
r.Post("/auth/register", api.Register)
|
||||
r.Post("/auth/login", api.Login)
|
||||
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(middleware.RequireBearer(authSvc))
|
||||
|
||||
r.Post("/auth/logout", api.Logout)
|
||||
r.Get("/me", api.Me)
|
||||
r.Get("/me/tokens", api.ListTokens)
|
||||
r.Post("/me/tokens", api.CreateToken)
|
||||
r.Delete("/me/tokens/{id}", api.RevokeToken)
|
||||
|
||||
// repositories (read for visible, write for permitted)
|
||||
r.Get("/repositories", api.ListRepositories)
|
||||
r.Post("/repositories", api.CreateRepository)
|
||||
r.Get("/repositories/{repo}", api.GetRepository)
|
||||
r.Patch("/repositories/{repo}", api.UpdateRepository)
|
||||
r.Delete("/repositories/{repo}", api.DeleteRepository)
|
||||
r.Get("/repositories/{repo}/members", api.ListMembers)
|
||||
r.Post("/repositories/{repo}/members", api.AddMember)
|
||||
r.Patch("/repositories/{repo}/members/{username}", api.UpdateMember)
|
||||
r.Delete("/repositories/{repo}/members/{username}", api.RemoveMember)
|
||||
r.Get("/repositories/{repo}/pubkey", api.RepoPubkey)
|
||||
|
||||
// structure
|
||||
r.Get("/repositories/{repo}/distributions", api.ListDistributions)
|
||||
r.Post("/repositories/{repo}/distributions", api.CreateDistribution)
|
||||
r.Delete("/repositories/{repo}/distributions/{dist}", api.DeleteDistribution)
|
||||
r.Get("/repositories/{repo}/distributions/{dist}/components", api.ListComponents)
|
||||
r.Post("/repositories/{repo}/distributions/{dist}/components", api.CreateComponent)
|
||||
r.Delete("/repositories/{repo}/distributions/{dist}/components/{comp}", api.DeleteComponent)
|
||||
r.Get("/repositories/{repo}/distributions/{dist}/architectures", api.ListArchitectures)
|
||||
r.Post("/repositories/{repo}/distributions/{dist}/architectures", api.CreateArchitecture)
|
||||
r.Delete("/repositories/{repo}/distributions/{dist}/architectures/{arch}", api.DeleteArchitecture)
|
||||
|
||||
// packages
|
||||
r.Get("/repositories/{repo}/distributions/{dist}/packages", api.ListPackages)
|
||||
r.Post("/repositories/{repo}/distributions/{dist}/packages", api.PushPackage)
|
||||
r.Get("/repositories/{repo}/packages/{id}", api.GetPackage)
|
||||
r.Get("/repositories/{repo}/packages/{id}/file", api.GetPackageFile)
|
||||
r.Delete("/repositories/{repo}/packages/{id}", api.DeletePackage)
|
||||
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(api.RequireAdmin)
|
||||
r.Get("/users", api.ListUsers)
|
||||
r.Get("/users/{id}", api.GetUser)
|
||||
r.Patch("/users/{id}", api.UpdateUser)
|
||||
r.Delete("/users/{id}", api.DeleteUser)
|
||||
})
|
||||
})
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
// RequireAdmin is middleware that requires the caller to be a server admin.
|
||||
func (api *API) RequireAdmin(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
id := middleware.IdentityFromContext(r.Context())
|
||||
if id == nil || !id.User.IsAdmin {
|
||||
httputil.WriteError(w, http.StatusForbidden, httputil.CodeForbidden, "admin privileges required")
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
// validation patterns.
|
||||
var (
|
||||
usernameRE = regexp.MustCompile(`^[a-z0-9_-]{3,32}$`)
|
||||
passwordRE = regexp.MustCompile(`^.{8,256}$`)
|
||||
tokenNameRE = regexp.MustCompile(`^.{1,64}$`)
|
||||
)
|
||||
|
||||
// validateUsername returns true if s is an acceptable username.
|
||||
func validateUsername(s string) bool { return usernameRE.MatchString(s) }
|
||||
|
||||
// validatePassword returns true if s is an acceptable password.
|
||||
func validatePassword(s string) bool { return passwordRE.MatchString(s) }
|
||||
|
||||
// bad renders a 400 validation error.
|
||||
func bad(w http.ResponseWriter, msg string) {
|
||||
httputil.WriteError(w, http.StatusBadRequest, httputil.CodeBadRequest, msg)
|
||||
}
|
||||
|
||||
// contextKey for request-scoped values is not needed beyond middleware; this
|
||||
// var keeps context imported if future handlers need it.
|
||||
var _ = context.Background
|
||||
@@ -0,0 +1,263 @@
|
||||
package restapi
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
|
||||
"urapt/server/auth"
|
||||
"urapt/server/cache"
|
||||
"urapt/server/store"
|
||||
"urapt/shared/config"
|
||||
"urapt/shared/db"
|
||||
"urapt/shared/gpg"
|
||||
)
|
||||
|
||||
type harness struct {
|
||||
t *testing.T
|
||||
srv *httptest.Server
|
||||
store *store.Store
|
||||
token string
|
||||
pkgDir string
|
||||
}
|
||||
|
||||
func newHarness(t *testing.T) *harness {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
database, err := db.Open(filepath.Join(dir, "test.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("db open: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { database.Close() })
|
||||
st := store.New(database)
|
||||
authSvc := auth.NewService(st)
|
||||
|
||||
key, err := gpg.GenerateKey("urapt-test <test>", 2048)
|
||||
if err != nil {
|
||||
t.Fatalf("gpg key: %v", err)
|
||||
}
|
||||
sp := &testSigner{key: key}
|
||||
cfg := config.Defaults
|
||||
cfg.StoreDir = dir
|
||||
cfg.PackagesDir = filepath.Join(dir, "packages")
|
||||
cfg.DBPath = filepath.Join(dir, "test.db")
|
||||
if err := os.MkdirAll(cfg.PackagesDir, 0o755); err != nil {
|
||||
t.Fatalf("mkdir pkg dir: %v", err)
|
||||
}
|
||||
|
||||
h := &harness{t: t, store: st, pkgDir: cfg.PackagesDir}
|
||||
mux := New(st, authSvc, sp, &cfg, cache.New())
|
||||
root := chi.NewRouter()
|
||||
root.Mount("/api/v1", mux)
|
||||
h.srv = httptest.NewServer(root)
|
||||
t.Cleanup(h.srv.Close)
|
||||
return h
|
||||
}
|
||||
|
||||
func (h *harness) do(method, path, token string, body any) (int, []byte) {
|
||||
h.t.Helper()
|
||||
var r io.Reader
|
||||
if body != nil {
|
||||
b, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
h.t.Fatalf("marshal: %v", err)
|
||||
}
|
||||
r = bytes.NewReader(b)
|
||||
}
|
||||
req, err := http.NewRequest(method, h.srv.URL+path, r)
|
||||
if err != nil {
|
||||
h.t.Fatalf("new req: %v", err)
|
||||
}
|
||||
if body != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
if token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
}
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
h.t.Fatalf("do %s %s: %v", method, path, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
data, _ := io.ReadAll(resp.Body)
|
||||
return resp.StatusCode, data
|
||||
}
|
||||
|
||||
func (h *harness) serverInfo() (int, map[string]any) {
|
||||
code, body := h.do("GET", "/api/v1/server/info", "", nil)
|
||||
var m map[string]any
|
||||
_ = json.Unmarshal(body, &m)
|
||||
return code, m
|
||||
}
|
||||
|
||||
func TestServerInfoNeedsSetup(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
code, m := h.serverInfo()
|
||||
if code != 200 {
|
||||
t.Fatalf("status %d", code)
|
||||
}
|
||||
if m["needs_setup"] != true {
|
||||
t.Fatalf("expected needs_setup=true, got %v", m["needs_setup"])
|
||||
}
|
||||
if m["default_key_fingerprint"] == "" {
|
||||
t.Fatal("expected fingerprint")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterFirstUserIsAdmin(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
code, body := h.do("POST", "/api/v1/auth/register", "", map[string]string{
|
||||
"username": "alice", "password": "supersecret",
|
||||
})
|
||||
if code != 201 {
|
||||
t.Fatalf("register status %d body %s", code, body)
|
||||
}
|
||||
var resp struct {
|
||||
User map[string]any `json:"user"`
|
||||
Token string `json:"token"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &resp); err != nil {
|
||||
t.Fatalf("unmarshal: %v", err)
|
||||
}
|
||||
if resp.Token == "" {
|
||||
t.Fatal("empty token")
|
||||
}
|
||||
if resp.User["is_admin"] != true {
|
||||
t.Fatalf("first user should be admin, got %v", resp.User["is_admin"])
|
||||
}
|
||||
h.token = resp.Token
|
||||
|
||||
// /me with token
|
||||
code, body = h.do("GET", "/api/v1/me", h.token, nil)
|
||||
if code != 200 {
|
||||
t.Fatalf("me status %d", code)
|
||||
}
|
||||
|
||||
// needs_setup should now be false
|
||||
_, m := h.serverInfo()
|
||||
if m["needs_setup"] != false {
|
||||
t.Fatalf("expected needs_setup=false after register")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterRejectsBadUsername(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
code, _ := h.do("POST", "/api/v1/auth/register", "", map[string]string{
|
||||
"username": "A", "password": "supersecret",
|
||||
})
|
||||
if code != 400 {
|
||||
t.Fatalf("expected 400 for short username, got %d", code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginAndAuthFlow(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
h.do("POST", "/api/v1/auth/register", "", map[string]string{
|
||||
"username": "bob", "password": "supersecret",
|
||||
})
|
||||
|
||||
code, body := h.do("POST", "/api/v1/auth/login", "", map[string]string{
|
||||
"username": "bob", "password": "supersecret",
|
||||
})
|
||||
if code != 200 {
|
||||
t.Fatalf("login status %d", code)
|
||||
}
|
||||
var resp struct {
|
||||
Token string `json:"token"`
|
||||
}
|
||||
json.Unmarshal(body, &resp)
|
||||
if resp.Token == "" {
|
||||
t.Fatal("empty token")
|
||||
}
|
||||
|
||||
// wrong password
|
||||
code, _ = h.do("POST", "/api/v1/auth/login", "", map[string]string{
|
||||
"username": "bob", "password": "wrongpassword",
|
||||
})
|
||||
if code != 401 {
|
||||
t.Fatalf("expected 401 for wrong password, got %d", code)
|
||||
}
|
||||
|
||||
// me without token
|
||||
code, _ = h.do("GET", "/api/v1/me", "", nil)
|
||||
if code != 401 {
|
||||
t.Fatalf("expected 401 without token, got %d", code)
|
||||
}
|
||||
|
||||
// tokens list
|
||||
code, _ = h.do("GET", "/api/v1/me/tokens", resp.Token, nil)
|
||||
if code != 200 {
|
||||
t.Fatalf("tokens list status %d", code)
|
||||
}
|
||||
|
||||
// create token
|
||||
code, body = h.do("POST", "/api/v1/me/tokens", resp.Token, map[string]string{"name": "laptop"})
|
||||
if code != 201 {
|
||||
t.Fatalf("create token status %d", code)
|
||||
}
|
||||
|
||||
// logout
|
||||
code, _ = h.do("POST", "/api/v1/auth/logout", resp.Token, nil)
|
||||
if code != 204 {
|
||||
t.Fatalf("logout status %d", code)
|
||||
}
|
||||
// token now revoked
|
||||
code, _ = h.do("GET", "/api/v1/me", resp.Token, nil)
|
||||
if code != 401 {
|
||||
t.Fatalf("expected 401 after logout, got %d", code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUsersAdminOnly(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
// register two users; first is admin
|
||||
_, body := h.do("POST", "/api/v1/auth/register", "", map[string]string{"username": "admin", "password": "supersecret"})
|
||||
var admin struct {
|
||||
Token string `json:"token"`
|
||||
}
|
||||
json.Unmarshal(body, &admin)
|
||||
|
||||
_, body = h.do("POST", "/api/v1/auth/register", "", map[string]string{"username": "carol", "password": "supersecret"})
|
||||
var carol struct {
|
||||
User map[string]any `json:"user"`
|
||||
Token string `json:"token"`
|
||||
}
|
||||
json.Unmarshal(body, &carol)
|
||||
if carol.User["is_admin"] == true {
|
||||
t.Fatal("second user should not be admin")
|
||||
}
|
||||
|
||||
// carol cannot list users
|
||||
code, _ := h.do("GET", "/api/v1/users", carol.Token, nil)
|
||||
if code != 403 {
|
||||
t.Fatalf("non-admin list users should be 403, got %d", code)
|
||||
}
|
||||
// admin can list users
|
||||
code, body = h.do("GET", "/api/v1/users", admin.Token, nil)
|
||||
if code != 200 {
|
||||
t.Fatalf("admin list users should be 200, got %d", code)
|
||||
}
|
||||
|
||||
// admin cannot delete self
|
||||
code, _ = h.do("DELETE", "/api/v1/users/"+carol.User["id"].(string), admin.Token, nil)
|
||||
if code != 204 {
|
||||
t.Fatalf("admin delete carol should be 204, got %d", code)
|
||||
}
|
||||
}
|
||||
|
||||
// keep context import used in case of future expansion
|
||||
var _ = context.Background
|
||||
|
||||
// testSigner implements restapi.SignerProvider for tests.
|
||||
type testSigner struct{ key *gpg.Key }
|
||||
|
||||
func (s *testSigner) PublicKeyArmored() (string, error) { return s.key.ArmoredPublic() }
|
||||
func (s *testSigner) Fingerprint() string { return s.key.Fingerprint }
|
||||
@@ -0,0 +1,195 @@
|
||||
package restapi
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"urapt/server/middleware"
|
||||
"urapt/server/store"
|
||||
apitypes "urapt/shared/api"
|
||||
"urapt/shared/crypto"
|
||||
"urapt/shared/httputil"
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// Register creates a new account. The first account becomes the admin. When
|
||||
// open_registration is false and an account already exists, registration is
|
||||
// closed to non-admins.
|
||||
func (api *API) Register(w http.ResponseWriter, r *http.Request) {
|
||||
var req apitypes.RegisterRequest
|
||||
if err := httputil.ReadJSON(r, &req, 1<<20); err != nil {
|
||||
bad(w, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
req.Username = strings.TrimSpace(req.Username)
|
||||
if !validateUsername(req.Username) {
|
||||
bad(w, "username must be 3-32 chars of [a-z0-9_-]")
|
||||
return
|
||||
}
|
||||
if !validatePassword(req.Password) {
|
||||
bad(w, "password must be 8-256 chars")
|
||||
return
|
||||
}
|
||||
|
||||
users, err := api.Store.ListUsers(r.Context())
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to read users")
|
||||
return
|
||||
}
|
||||
if len(users) > 0 && !api.Config.OpenRegistration {
|
||||
httputil.WriteError(w, http.StatusForbidden, httputil.CodeForbidden, "registration is closed")
|
||||
return
|
||||
}
|
||||
|
||||
if _, err := api.Store.GetUserByUsername(r.Context(), req.Username); err == nil {
|
||||
httputil.WriteError(w, http.StatusConflict, httputil.CodeConflict, "username already taken")
|
||||
return
|
||||
} else if !errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to check username")
|
||||
return
|
||||
}
|
||||
|
||||
hash, err := crypto.HashPassword(req.Password)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to hash password")
|
||||
return
|
||||
}
|
||||
user, _, err := api.Store.CreateUser(r.Context(), req.Username, hash)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to create user")
|
||||
return
|
||||
}
|
||||
|
||||
token, err := api.issueToken(r.Context(), user.ID, "login")
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to issue token")
|
||||
return
|
||||
}
|
||||
_ = api.Store.RecordAudit(r.Context(), &user.ID, nil, "user.register", user.Username, "")
|
||||
httputil.WriteJSON(w, http.StatusCreated, apitypes.AuthResponse{User: user, Token: token})
|
||||
}
|
||||
|
||||
// Login authenticates a user and issues a new API token.
|
||||
func (api *API) Login(w http.ResponseWriter, r *http.Request) {
|
||||
var req apitypes.LoginRequest
|
||||
if err := httputil.ReadJSON(r, &req, 1<<20); err != nil {
|
||||
bad(w, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
user, err := api.Store.GetUserByUsername(r.Context(), req.Username)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusUnauthorized, httputil.CodeUnauthorized, "invalid credentials")
|
||||
return
|
||||
}
|
||||
// Need the password hash; fetch via a dedicated method.
|
||||
hash, err := api.Store.GetUserPasswordHash(r.Context(), user.ID)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to read user")
|
||||
return
|
||||
}
|
||||
if !crypto.VerifyPassword(hash, req.Password) {
|
||||
httputil.WriteError(w, http.StatusUnauthorized, httputil.CodeUnauthorized, "invalid credentials")
|
||||
return
|
||||
}
|
||||
token, err := api.issueToken(r.Context(), user.ID, "login")
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to issue token")
|
||||
return
|
||||
}
|
||||
_ = api.Store.RecordAudit(r.Context(), &user.ID, nil, "user.login", user.Username, "")
|
||||
httputil.WriteJSON(w, http.StatusOK, apitypes.AuthResponse{User: user, Token: token})
|
||||
}
|
||||
|
||||
// Logout revokes the caller's current token.
|
||||
func (api *API) Logout(w http.ResponseWriter, r *http.Request) {
|
||||
id := middleware.IdentityFromContext(r.Context())
|
||||
if err := api.Store.RevokeToken(r.Context(), id.User.ID, id.TokenID); err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to revoke token")
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
// Me returns the caller's user record.
|
||||
func (api *API) Me(w http.ResponseWriter, r *http.Request) {
|
||||
id := middleware.IdentityFromContext(r.Context())
|
||||
httputil.WriteJSON(w, http.StatusOK, id.User)
|
||||
}
|
||||
|
||||
// ListTokens returns the caller's tokens.
|
||||
func (api *API) ListTokens(w http.ResponseWriter, r *http.Request) {
|
||||
id := middleware.IdentityFromContext(r.Context())
|
||||
tokens, err := api.Store.ListTokens(r.Context(), id.User.ID)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to list tokens")
|
||||
return
|
||||
}
|
||||
if tokens == nil {
|
||||
tokens = []*models.APIToken{}
|
||||
}
|
||||
httputil.WriteJSON(w, http.StatusOK, apitypes.ListResponse[*models.APIToken]{Items: tokens, Page: 1, PerPage: 100, Total: len(tokens)})
|
||||
}
|
||||
|
||||
// CreateToken issues a new named token for the caller.
|
||||
func (api *API) CreateToken(w http.ResponseWriter, r *http.Request) {
|
||||
var req apitypes.CreateTokenRequest
|
||||
if err := httputil.ReadJSON(r, &req, 1<<20); err != nil {
|
||||
bad(w, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
if !tokenNameRE.MatchString(req.Name) {
|
||||
bad(w, "name must be 1-64 chars")
|
||||
return
|
||||
}
|
||||
id := middleware.IdentityFromContext(r.Context())
|
||||
token, row, err := api.issueTokenRow(r.Context(), id.User.ID, req.Name)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to issue token")
|
||||
return
|
||||
}
|
||||
row.Token = token
|
||||
_ = api.Store.RecordAudit(r.Context(), &id.User.ID, nil, "token.create", req.Name, "")
|
||||
httputil.WriteJSON(w, http.StatusCreated, row)
|
||||
}
|
||||
|
||||
// RevokeToken revokes one of the caller's tokens by id.
|
||||
func (api *API) RevokeToken(w http.ResponseWriter, r *http.Request) {
|
||||
id := middleware.IdentityFromContext(r.Context())
|
||||
tokenID := r.PathValue("id")
|
||||
if err := api.Store.RevokeToken(r.Context(), id.User.ID, tokenID); err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "token not found")
|
||||
return
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to revoke token")
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
// issueToken generates a token, persists its hash, and returns the plaintext.
|
||||
func (api *API) issueToken(ctx context.Context, userID, name string) (string, error) {
|
||||
token, hash, prefix, err := crypto.GenerateToken()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if _, err := api.Store.CreateToken(ctx, userID, name, prefix, hash); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
// issueTokenRow is like issueToken but also returns the persisted token row.
|
||||
func (api *API) issueTokenRow(ctx context.Context, userID, name string) (string, *models.APIToken, error) {
|
||||
token, hash, prefix, err := crypto.GenerateToken()
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
row, err := api.Store.CreateToken(ctx, userID, name, prefix, hash)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
return token, row, nil
|
||||
}
|
||||
@@ -0,0 +1,339 @@
|
||||
package restapi
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"urapt/server/middleware"
|
||||
"urapt/server/store"
|
||||
"urapt/shared/apt"
|
||||
"urapt/shared/deb"
|
||||
"urapt/shared/httputil"
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// ListPackages lists packages in a (repo, distribution) with optional filters.
|
||||
func (api *API) ListPackages(w http.ResponseWriter, r *http.Request) {
|
||||
repo, dist := api.loadDistro(w, r)
|
||||
if dist == nil {
|
||||
return
|
||||
}
|
||||
f := store.PackageFilters{
|
||||
ComponentID: r.URL.Query().Get("component"),
|
||||
Arch: r.URL.Query().Get("arch"),
|
||||
Name: r.URL.Query().Get("name"),
|
||||
Query: r.URL.Query().Get("q"),
|
||||
}
|
||||
page, _ := strconv.Atoi(r.URL.Query().Get("page"))
|
||||
perPage, _ := strconv.Atoi(r.URL.Query().Get("per_page"))
|
||||
pkgs, total, err := api.Store.ListPackages(r.Context(), repo.ID, dist.ID, f, page, perPage)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to list packages")
|
||||
return
|
||||
}
|
||||
if pkgs == nil {
|
||||
pkgs = []*models.Package{}
|
||||
}
|
||||
httputil.WriteJSON(w, http.StatusOK, map[string]any{
|
||||
"items": pkgs, "page": pageOr(page), "per_page": perPageOr(perPage), "total": total,
|
||||
})
|
||||
}
|
||||
|
||||
// PushPackage receives a multipart .deb upload, validates it, stores the blob,
|
||||
// and records the package.
|
||||
func (api *API) PushPackage(w http.ResponseWriter, r *http.Request) {
|
||||
repo, dist := api.requireDistroWrite(w, r)
|
||||
if dist == nil {
|
||||
return
|
||||
}
|
||||
id := middleware.IdentityFromContext(r.Context())
|
||||
|
||||
// Stream the multipart upload to a temp file in the packages dir.
|
||||
tempPath, origName, componentName, err := api.receiveUpload(r)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusBadRequest, httputil.CodeBadRequest, err.Error())
|
||||
return
|
||||
}
|
||||
cleanup := func() { _ = os.Remove(tempPath) }
|
||||
defer func() { _ = os.Remove(tempPath) }()
|
||||
|
||||
if componentName == "" {
|
||||
bad(w, "component field is required")
|
||||
return
|
||||
}
|
||||
component, err := api.Store.GetComponentByName(r.Context(), dist.ID, componentName)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusBadRequest, httputil.CodeBadRequest, "component not found in distribution")
|
||||
return
|
||||
}
|
||||
|
||||
// Parse and hash the uploaded .deb.
|
||||
inspected, err := deb.Inspect(tempPath)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusBadRequest, httputil.CodeBadRequest, "invalid .deb: "+err.Error())
|
||||
return
|
||||
}
|
||||
ctrl := inspected.Control
|
||||
if ctrl.Get("Package") == "" || ctrl.Get("Version") == "" || ctrl.Get("Architecture") == "" {
|
||||
httputil.WriteError(w, http.StatusBadRequest, httputil.CodeBadRequest, "control missing Package/Version/Architecture")
|
||||
return
|
||||
}
|
||||
|
||||
// Validate architecture is configured (or "all").
|
||||
pkgArch := ctrl.Get("Architecture")
|
||||
if pkgArch != "all" {
|
||||
ok, err := api.Store.HasArchitecture(r.Context(), dist.ID, pkgArch)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to check architecture")
|
||||
return
|
||||
}
|
||||
if !ok {
|
||||
httputil.WriteError(w, http.StatusBadRequest, httputil.CodeBadRequest, "architecture "+pkgArch+" not configured for distribution")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Dedup on (repo, distro, component, name, version, arch).
|
||||
existing, err := api.Store.GetPackageByPoolPath(r.Context(), repo.ID,
|
||||
apt.PoolPath(component.Name, ctrl.Get("Source"), ctrl.Get("Package"), origName))
|
||||
_ = existing
|
||||
if err == nil {
|
||||
httputil.WriteError(w, http.StatusConflict, httputil.CodeConflict,
|
||||
"package "+ctrl.Get("Package")+"_"+ctrl.Get("Version")+"_"+pkgArch+" already exists")
|
||||
return
|
||||
} else if !errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to check duplicate")
|
||||
return
|
||||
}
|
||||
|
||||
// Find-or-create the content-addressed blob.
|
||||
blobFileName := api.blobFilePath(inspected.SHA256)
|
||||
created, err := api.Store.CreateBlob(r.Context(), inspected.SHA256, blobFileName, inspected.Size)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to record blob")
|
||||
return
|
||||
}
|
||||
if created {
|
||||
// New blob: move the temp file into place.
|
||||
if err := os.Rename(tempPath, blobFileName); err != nil {
|
||||
_ = api.Store.DeleteBlob(r.Context(), inspected.SHA256)
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to store package file")
|
||||
return
|
||||
}
|
||||
} else {
|
||||
// Existing blob: increment ref count and discard the temp upload.
|
||||
if _, err := api.Store.IncBlobRef(r.Context(), inspected.SHA256); err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to increment blob ref")
|
||||
return
|
||||
}
|
||||
cleanup()
|
||||
}
|
||||
|
||||
pool := apt.PoolPath(component.Name, ctrl.Get("Source"), ctrl.Get("Package"), origName)
|
||||
pkg := &models.Package{
|
||||
RepositoryID: repo.ID,
|
||||
DistributionID: dist.ID,
|
||||
ComponentID: component.ID,
|
||||
Name: ctrl.Get("Package"),
|
||||
Version: ctrl.Get("Version"),
|
||||
Architecture: pkgArch,
|
||||
Source: ctrl.Get("Source"),
|
||||
Maintainer: ctrl.Get("Maintainer"),
|
||||
Priority: ctrl.Get("Priority"),
|
||||
Section: ctrl.Get("Section"),
|
||||
Origin: ctrl.Get("Origin"),
|
||||
Homepage: ctrl.Get("Homepage"),
|
||||
Description: ctrl.Get("Description"),
|
||||
DescriptionMD5: ctrl.DescriptionMD5(),
|
||||
Depends: ctrl.Get("Depends"),
|
||||
PreDepends: ctrl.Get("Pre-Depends"),
|
||||
Recommends: ctrl.Get("Recommends"),
|
||||
Suggests: ctrl.Get("Suggests"),
|
||||
Conflicts: ctrl.Get("Conflicts"),
|
||||
Breaks: ctrl.Get("Breaks"),
|
||||
Provides: ctrl.Get("Provides"),
|
||||
Replaces: ctrl.Get("Replaces"),
|
||||
Enhances: ctrl.Get("Enhances"),
|
||||
InstalledSize: parseInt64(ctrl.Get("Installed-Size")),
|
||||
Essential: ctrl.Get("Essential"),
|
||||
BuiltUsing: ctrl.Get("Built-Using"),
|
||||
Tag: ctrl.Get("Tag"),
|
||||
RawControl: ctrl.Raw,
|
||||
Filename: blobFileName,
|
||||
PoolPath: pool,
|
||||
Size: inspected.Size,
|
||||
MD5sum: inspected.MD5sum,
|
||||
SHA1: inspected.SHA1,
|
||||
SHA256: inspected.SHA256,
|
||||
UploadedByUserID: id.User.ID,
|
||||
}
|
||||
if err := api.Store.CreatePackage(r.Context(), pkg); err != nil {
|
||||
// Roll back the blob ref we added.
|
||||
if rc, _ := api.Store.DecBlobRef(r.Context(), inspected.SHA256); rc == 0 {
|
||||
_ = api.Store.DeleteBlob(r.Context(), inspected.SHA256)
|
||||
_ = os.Remove(blobFileName)
|
||||
}
|
||||
httputil.WriteError(w, http.StatusConflict, httputil.CodeConflict, "package already exists or invalid")
|
||||
return
|
||||
}
|
||||
api.Cache.Invalidate(repo.ID, dist.Name)
|
||||
_ = api.Store.RecordAudit(r.Context(), &id.User.ID, &repo.ID, "package.push", pkg.Name+"_"+pkg.Version, pkg.Architecture)
|
||||
httputil.WriteJSON(w, http.StatusCreated, pkg)
|
||||
}
|
||||
|
||||
// GetPackage returns a single package by id.
|
||||
func (api *API) GetPackage(w http.ResponseWriter, r *http.Request) {
|
||||
repo := api.requireRead(w, r)
|
||||
if repo == nil {
|
||||
return
|
||||
}
|
||||
pkg, err := api.Store.GetPackageByID(r.Context(), r.PathValue("id"))
|
||||
if err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "package not found")
|
||||
return
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to read package")
|
||||
return
|
||||
}
|
||||
if pkg.RepositoryID != repo.ID {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "package not found")
|
||||
return
|
||||
}
|
||||
httputil.WriteJSON(w, http.StatusOK, pkg)
|
||||
}
|
||||
|
||||
// GetPackageFile streams a package's .deb file (used by the CLI pull command).
|
||||
func (api *API) GetPackageFile(w http.ResponseWriter, r *http.Request) {
|
||||
repo := api.requireRead(w, r)
|
||||
if repo == nil {
|
||||
return
|
||||
}
|
||||
pkg, err := api.Store.GetPackageByID(r.Context(), r.PathValue("id"))
|
||||
if err != nil || pkg.RepositoryID != repo.ID {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "package not found")
|
||||
return
|
||||
}
|
||||
path := api.blobFilePath(pkg.SHA256)
|
||||
w.Header().Set("Content-Disposition", `attachment; filename="`+filepath.Base(pkg.PoolPath)+`"`)
|
||||
http.ServeFile(w, r, path)
|
||||
}
|
||||
|
||||
// DeletePackage removes a package and decrements its blob reference.
|
||||
func (api *API) DeletePackage(w http.ResponseWriter, r *http.Request) {
|
||||
repo := api.requireWrite(w, r)
|
||||
if repo == nil {
|
||||
return
|
||||
}
|
||||
pkg, err := api.Store.GetPackageByID(r.Context(), r.PathValue("id"))
|
||||
if err != nil || pkg.RepositoryID != repo.ID {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "package not found")
|
||||
return
|
||||
}
|
||||
if err := api.Store.DeletePackage(r.Context(), pkg.ID); err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to delete package")
|
||||
return
|
||||
}
|
||||
if rc, _ := api.Store.DecBlobRef(r.Context(), pkg.SHA256); rc == 0 {
|
||||
_ = api.Store.DeleteBlob(r.Context(), pkg.SHA256)
|
||||
_ = os.Remove(api.blobFilePath(pkg.SHA256))
|
||||
}
|
||||
api.Cache.InvalidateRepo(repo.ID)
|
||||
id := middleware.IdentityFromContext(r.Context())
|
||||
_ = api.Store.RecordAudit(r.Context(), &id.User.ID, &repo.ID, "package.delete", pkg.Name+"_"+pkg.Version, pkg.Architecture)
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
// receiveUpload streams a multipart upload (field "file" + "component") to a
|
||||
// temp file in the packages directory, enforcing the size limit. Returns the
|
||||
// temp path, original filename, component name, and any error.
|
||||
func (api *API) receiveUpload(r *http.Request) (tempPath, origName, component string, err error) {
|
||||
reader, err := r.MultipartReader()
|
||||
if err != nil {
|
||||
return "", "", "", errors.New("expected multipart/form-data")
|
||||
}
|
||||
maxSize := api.Config.MaxPackageSize
|
||||
|
||||
f, err := os.CreateTemp(api.Config.PackagesDir, ".upload-*")
|
||||
if err != nil {
|
||||
return "", "", "", errors.New("failed to create temp file")
|
||||
}
|
||||
tempPath = f.Name()
|
||||
defer func() { _ = f.Close() }()
|
||||
|
||||
gotFile := false
|
||||
gotComponent := false
|
||||
var written int64
|
||||
for {
|
||||
part, perr := reader.NextPart()
|
||||
if perr == io.EOF {
|
||||
break
|
||||
}
|
||||
if perr != nil {
|
||||
return tempPath, "", "", perr
|
||||
}
|
||||
switch part.FormName() {
|
||||
case "component":
|
||||
data, derr := io.ReadAll(io.LimitReader(part, 256))
|
||||
if derr != nil {
|
||||
return tempPath, "", "", derr
|
||||
}
|
||||
component = strings.TrimSpace(string(data))
|
||||
gotComponent = true
|
||||
case "file":
|
||||
origName = part.FileName()
|
||||
if origName == "" {
|
||||
return tempPath, "", "", errors.New("file field has no filename")
|
||||
}
|
||||
n, werr := io.Copy(f, io.LimitReader(part, maxSize+1))
|
||||
if werr != nil {
|
||||
return tempPath, origName, "", werr
|
||||
}
|
||||
written = n
|
||||
gotFile = true
|
||||
default:
|
||||
// ignore unknown fields
|
||||
}
|
||||
}
|
||||
if !gotFile {
|
||||
return tempPath, "", "", errors.New("missing 'file' field")
|
||||
}
|
||||
if !gotComponent {
|
||||
return tempPath, origName, "", errors.New("missing 'component' field")
|
||||
}
|
||||
if written > maxSize {
|
||||
return tempPath, origName, "", errors.New("package exceeds max size")
|
||||
}
|
||||
_ = gotComponent
|
||||
return tempPath, origName, component, nil
|
||||
}
|
||||
|
||||
// parseInt64 parses a base-10 int64, returning 0 on error.
|
||||
func parseInt64(s string) int64 {
|
||||
n, _ := strconv.ParseInt(s, 10, 64)
|
||||
return n
|
||||
}
|
||||
|
||||
// pageOr defaults page to 1.
|
||||
func pageOr(p int) int {
|
||||
if p < 1 {
|
||||
return 1
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
// perPageOr defaults per-page to 25.
|
||||
func perPageOr(p int) int {
|
||||
if p < 1 {
|
||||
return 25
|
||||
}
|
||||
if p > 100 {
|
||||
return 100
|
||||
}
|
||||
return p
|
||||
}
|
||||
@@ -0,0 +1,192 @@
|
||||
package restapi
|
||||
|
||||
import (
|
||||
"archive/tar"
|
||||
"bytes"
|
||||
"compress/gzip"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// buildDebBytes constructs a minimal valid .deb with the given control text.
|
||||
func buildDebBytes(t *testing.T, controlText string) []byte {
|
||||
t.Helper()
|
||||
var ctrlBuf bytes.Buffer
|
||||
gz := gzip.NewWriter(&ctrlBuf)
|
||||
tw := tar.NewWriter(gz)
|
||||
writeTarFile(t, tw, "control", controlText)
|
||||
tw.Close()
|
||||
gz.Close()
|
||||
|
||||
var dataBuf bytes.Buffer
|
||||
gz2 := gzip.NewWriter(&dataBuf)
|
||||
tw2 := tar.NewWriter(gz2)
|
||||
writeTarFile(t, tw2, "usr/share/foo", "x")
|
||||
tw2.Close()
|
||||
gz2.Close()
|
||||
|
||||
var out bytes.Buffer
|
||||
out.WriteString("!<arch>\n")
|
||||
writeArMemberBytes(&out, "debian-binary", []byte("2.0\n"))
|
||||
writeArMemberBytes(&out, "control.tar.gz", ctrlBuf.Bytes())
|
||||
writeArMemberBytes(&out, "data.tar.gz", dataBuf.Bytes())
|
||||
return out.Bytes()
|
||||
}
|
||||
|
||||
func writeTarFile(t *testing.T, tw *tar.Writer, name, body string) {
|
||||
t.Helper()
|
||||
if err := tw.WriteHeader(&tar.Header{Name: name, Mode: 0o644, Size: int64(len(body)), Typeflag: tar.TypeReg}); err != nil {
|
||||
t.Fatalf("tar header: %v", err)
|
||||
}
|
||||
if _, err := tw.Write([]byte(body)); err != nil {
|
||||
t.Fatalf("tar write: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func writeArMemberBytes(buf *bytes.Buffer, name string, data []byte) {
|
||||
header := make([]byte, 60)
|
||||
for i := range header {
|
||||
header[i] = ' '
|
||||
}
|
||||
copy(header[0:], name+"/")
|
||||
ds := []byte(padLeftInt(len(data), 10))
|
||||
copy(header[48:], ds)
|
||||
header[58] = '`'
|
||||
header[59] = '\n'
|
||||
buf.Write(header)
|
||||
buf.Write(data)
|
||||
if len(data)%2 == 1 {
|
||||
buf.WriteByte('\n')
|
||||
}
|
||||
}
|
||||
|
||||
func padLeftInt(n, width int) string {
|
||||
s := make([]byte, width)
|
||||
for i := range s {
|
||||
s[i] = ' '
|
||||
}
|
||||
digits := []byte(itoaInt(n))
|
||||
copy(s[len(s)-len(digits):], digits)
|
||||
return string(s)
|
||||
}
|
||||
|
||||
func itoaInt(n int) string {
|
||||
if n == 0 {
|
||||
return "0"
|
||||
}
|
||||
var b []byte
|
||||
for n > 0 {
|
||||
b = append([]byte{byte('0' + n%10)}, b...)
|
||||
n /= 10
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
|
||||
// uploadPackage POSTs a multipart push.
|
||||
func (h *harness) uploadPackage(token, repo, dist, component string, deb []byte) (int, []byte) {
|
||||
h.t.Helper()
|
||||
var buf bytes.Buffer
|
||||
mw := multipart.NewWriter(&buf)
|
||||
_ = mw.WriteField("component", component)
|
||||
fw, _ := mw.CreateFormFile("file", "foo_1.0_amd64.deb")
|
||||
fw.Write(deb)
|
||||
mw.Close()
|
||||
|
||||
req, _ := http.NewRequest("POST", h.srv.URL+"/api/v1/repositories/"+repo+"/distributions/"+dist+"/packages", &buf)
|
||||
req.Header.Set("Content-Type", mw.FormDataContentType())
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
h.t.Fatalf("upload: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return resp.StatusCode, body
|
||||
}
|
||||
|
||||
func TestPushPullDeletePackage(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
_, body := h.do("POST", "/api/v1/auth/register", "", map[string]string{"username": "owner", "password": "supersecret"})
|
||||
var owner struct {
|
||||
Token string `json:"token"`
|
||||
}
|
||||
json.Unmarshal(body, &owner)
|
||||
|
||||
// Set packages dir so blobs land in the temp dir.
|
||||
_ = h.store // packages dir is taken from config (temp dir in harness).
|
||||
|
||||
h.do("POST", "/api/v1/repositories", owner.Token, map[string]any{"name": "pkgrepo", "visibility": "public"})
|
||||
h.do("POST", "/api/v1/repositories/pkgrepo/distributions", owner.Token, map[string]string{"name": "stable"})
|
||||
h.do("POST", "/api/v1/repositories/pkgrepo/distributions/stable/components", owner.Token, map[string]string{"name": "main"})
|
||||
h.do("POST", "/api/v1/repositories/pkgrepo/distributions/stable/architectures", owner.Token, map[string]string{"name": "amd64"})
|
||||
|
||||
debBytes := buildDebBytes(t, "Package: foo\nVersion: 1.0\nArchitecture: amd64\nMaintainer: Test <t@e.com>\nDescription: short\n extended\n")
|
||||
|
||||
code, body := h.uploadPackage(owner.Token, "pkgrepo", "stable", "main", debBytes)
|
||||
if code != 201 {
|
||||
t.Fatalf("push status %d body %s", code, body)
|
||||
}
|
||||
var pkg struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Pool string `json:"pool_path"`
|
||||
SHA string `json:"sha256"`
|
||||
}
|
||||
json.Unmarshal(body, &pkg)
|
||||
if pkg.Name != "foo" || pkg.SHA == "" {
|
||||
t.Fatalf("unexpected pkg: %+v", pkg)
|
||||
}
|
||||
|
||||
// duplicate push should 409
|
||||
code, _ = h.uploadPackage(owner.Token, "pkgrepo", "stable", "main", debBytes)
|
||||
if code != 409 {
|
||||
t.Fatalf("duplicate push should be 409, got %d", code)
|
||||
}
|
||||
|
||||
// list
|
||||
code, body = h.do("GET", "/api/v1/repositories/pkgrepo/distributions/stable/packages", owner.Token, nil)
|
||||
if code != 200 {
|
||||
t.Fatalf("list packages status %d", code)
|
||||
}
|
||||
|
||||
// get
|
||||
code, _ = h.do("GET", "/api/v1/repositories/pkgrepo/packages/"+pkg.ID, owner.Token, nil)
|
||||
if code != 200 {
|
||||
t.Fatalf("get package status %d", code)
|
||||
}
|
||||
|
||||
// download file
|
||||
code, body = h.do("GET", "/api/v1/repositories/pkgrepo/packages/"+pkg.ID+"/file", owner.Token, nil)
|
||||
if code != 200 {
|
||||
t.Fatalf("get file status %d", code)
|
||||
}
|
||||
if !bytes.Equal(body, debBytes) {
|
||||
t.Fatalf("downloaded file does not match uploaded (%d vs %d bytes)", len(body), len(debBytes))
|
||||
}
|
||||
|
||||
// verify blob file exists on disk in the temp packages dir
|
||||
blobPath := filepath.Join(h.pkgDir, pkg.SHA+".deb")
|
||||
if _, err := os.Stat(blobPath); err != nil {
|
||||
t.Fatalf("blob file missing on disk: %v", err)
|
||||
}
|
||||
|
||||
// delete
|
||||
code, _ = h.do("DELETE", "/api/v1/repositories/pkgrepo/packages/"+pkg.ID, owner.Token, nil)
|
||||
if code != 204 {
|
||||
t.Fatalf("delete package status %d", code)
|
||||
}
|
||||
// get now 404
|
||||
code, _ = h.do("GET", "/api/v1/repositories/pkgrepo/packages/"+pkg.ID, owner.Token, nil)
|
||||
if code != 404 {
|
||||
t.Fatalf("deleted package should 404, got %d", code)
|
||||
}
|
||||
// blob file removed after refcount hits 0
|
||||
if _, err := os.Stat(blobPath); !os.IsNotExist(err) {
|
||||
t.Fatalf("blob file should be removed after delete, err=%v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,368 @@
|
||||
package restapi
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"urapt/server/middleware"
|
||||
"urapt/server/store"
|
||||
apitypes "urapt/shared/api"
|
||||
"urapt/shared/httputil"
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// nameRE is the shared validator for repository, distribution, component, and
|
||||
// architecture names: lowercase, starting alphanumeric, allowing -+., length
|
||||
// 1-64.
|
||||
var nameRE = regexp.MustCompile(`^[a-z0-9][a-z0-9+.\-]{0,63}$`)
|
||||
|
||||
// loadRepoByName fetches a repository by its {repo} path param, rendering the
|
||||
// appropriate error. nil is returned only after an error has been written.
|
||||
func (api *API) loadRepoByName(w http.ResponseWriter, r *http.Request) *models.Repository {
|
||||
name := r.PathValue("repo")
|
||||
repo, err := api.Store.GetRepositoryByName(r.Context(), name)
|
||||
if err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "repository not found")
|
||||
return nil
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to read repository")
|
||||
return nil
|
||||
}
|
||||
return repo
|
||||
}
|
||||
|
||||
// requireRead loads the repo and checks CanRead; returns the repo or nil (after
|
||||
// writing an error).
|
||||
func (api *API) requireRead(w http.ResponseWriter, r *http.Request) *models.Repository {
|
||||
repo := api.loadRepoByName(w, r)
|
||||
if repo == nil {
|
||||
return nil
|
||||
}
|
||||
id := middleware.IdentityFromContext(r.Context())
|
||||
ok, err := api.Auth.CanRead(r.Context(), id.User, repo)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "permission check failed")
|
||||
return nil
|
||||
}
|
||||
if !ok {
|
||||
httputil.WriteError(w, http.StatusForbidden, httputil.CodeForbidden, "no read access")
|
||||
return nil
|
||||
}
|
||||
return repo
|
||||
}
|
||||
|
||||
// requireWrite loads the repo and checks CanWrite.
|
||||
func (api *API) requireWrite(w http.ResponseWriter, r *http.Request) *models.Repository {
|
||||
repo := api.loadRepoByName(w, r)
|
||||
if repo == nil {
|
||||
return nil
|
||||
}
|
||||
id := middleware.IdentityFromContext(r.Context())
|
||||
ok, err := api.Auth.CanWrite(r.Context(), id.User, repo)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "permission check failed")
|
||||
return nil
|
||||
}
|
||||
if !ok {
|
||||
httputil.WriteError(w, http.StatusForbidden, httputil.CodeForbidden, "no write access")
|
||||
return nil
|
||||
}
|
||||
return repo
|
||||
}
|
||||
|
||||
// requireManage loads the repo and checks CanManage.
|
||||
func (api *API) requireManage(w http.ResponseWriter, r *http.Request) *models.Repository {
|
||||
repo := api.loadRepoByName(w, r)
|
||||
if repo == nil {
|
||||
return nil
|
||||
}
|
||||
id := middleware.IdentityFromContext(r.Context())
|
||||
ok, err := api.Auth.CanManage(r.Context(), id.User, repo)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "permission check failed")
|
||||
return nil
|
||||
}
|
||||
if !ok {
|
||||
httputil.WriteError(w, http.StatusForbidden, httputil.CodeForbidden, "manage access required")
|
||||
return nil
|
||||
}
|
||||
return repo
|
||||
}
|
||||
|
||||
// ListRepositories returns repositories visible to the caller.
|
||||
func (api *API) ListRepositories(w http.ResponseWriter, r *http.Request) {
|
||||
id := middleware.IdentityFromContext(r.Context())
|
||||
repos, err := api.Store.ListReposVisible(r.Context(), id.User.ID)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to list repositories")
|
||||
return
|
||||
}
|
||||
if repos == nil {
|
||||
repos = []*models.Repository{}
|
||||
}
|
||||
httputil.WriteJSON(w, http.StatusOK, apitypes.ListResponse[*models.Repository]{Items: repos, Page: 1, PerPage: 100, Total: len(repos)})
|
||||
}
|
||||
|
||||
// CreateRepository creates a new repository owned by the caller.
|
||||
func (api *API) CreateRepository(w http.ResponseWriter, r *http.Request) {
|
||||
var req apitypes.CreateRepoRequest
|
||||
if err := httputil.ReadJSON(r, &req, 1<<20); err != nil {
|
||||
bad(w, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
req.Name = strings.TrimSpace(req.Name)
|
||||
if !nameRE.MatchString(req.Name) {
|
||||
bad(w, "name must be 1-64 chars of [a-z0-9][a-z0-9+.-]")
|
||||
return
|
||||
}
|
||||
vis := models.Visibility(strings.ToLower(req.Visibility))
|
||||
if vis != models.VisibilityPublic && vis != models.VisibilityPrivate {
|
||||
bad(w, "visibility must be 'public' or 'private'")
|
||||
return
|
||||
}
|
||||
id := middleware.IdentityFromContext(r.Context())
|
||||
if _, err := api.Store.GetRepositoryByName(r.Context(), req.Name); err == nil {
|
||||
httputil.WriteError(w, http.StatusConflict, httputil.CodeConflict, "repository name already taken")
|
||||
return
|
||||
} else if !errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to check name")
|
||||
return
|
||||
}
|
||||
repo, err := api.Store.CreateRepository(r.Context(), req.Name, id.User.ID, vis, req.Description)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to create repository")
|
||||
return
|
||||
}
|
||||
_ = api.Store.RecordAudit(r.Context(), &id.User.ID, &repo.ID, "repo.create", repo.Name, "")
|
||||
httputil.WriteJSON(w, http.StatusCreated, repo)
|
||||
}
|
||||
|
||||
// GetRepository returns a single repository.
|
||||
func (api *API) GetRepository(w http.ResponseWriter, r *http.Request) {
|
||||
repo := api.requireRead(w, r)
|
||||
if repo == nil {
|
||||
return
|
||||
}
|
||||
httputil.WriteJSON(w, http.StatusOK, repo)
|
||||
}
|
||||
|
||||
// UpdateRepository mutates a repository.
|
||||
func (api *API) UpdateRepository(w http.ResponseWriter, r *http.Request) {
|
||||
repo := api.requireManage(w, r)
|
||||
if repo == nil {
|
||||
return
|
||||
}
|
||||
var req apitypes.UpdateRepoRequest
|
||||
if err := httputil.ReadJSON(r, &req, 1<<20); err != nil {
|
||||
bad(w, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
var vis *models.Visibility
|
||||
if req.Visibility != nil {
|
||||
v := models.Visibility(strings.ToLower(*req.Visibility))
|
||||
if v != models.VisibilityPublic && v != models.VisibilityPrivate {
|
||||
bad(w, "visibility must be 'public' or 'private'")
|
||||
return
|
||||
}
|
||||
vis = &v
|
||||
}
|
||||
if req.Name != nil {
|
||||
if !nameRE.MatchString(*req.Name) {
|
||||
bad(w, "name must be 1-64 chars of [a-z0-9][a-z0-9+.-]")
|
||||
return
|
||||
}
|
||||
}
|
||||
if err := api.Store.UpdateRepository(r.Context(), repo.ID, valStr(req.Name), vis, req.Description); err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to update repository")
|
||||
return
|
||||
}
|
||||
api.Cache.InvalidateRepo(repo.ID)
|
||||
updated, _ := api.Store.GetRepositoryByID(r.Context(), repo.ID)
|
||||
httputil.WriteJSON(w, http.StatusOK, updated)
|
||||
}
|
||||
|
||||
// DeleteRepository removes a repository and cleans up its blobs.
|
||||
func (api *API) DeleteRepository(w http.ResponseWriter, r *http.Request) {
|
||||
repo := api.requireManage(w, r)
|
||||
if repo == nil {
|
||||
return
|
||||
}
|
||||
shas, err := api.Store.ListPackageBlobSHA256sByRepo(r.Context(), repo.ID)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to list packages")
|
||||
return
|
||||
}
|
||||
if err := api.Store.DeletePackagesByRepo(r.Context(), repo.ID); err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to delete packages")
|
||||
return
|
||||
}
|
||||
if err := api.Store.DeleteRepository(r.Context(), repo.ID); err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to delete repository")
|
||||
return
|
||||
}
|
||||
for _, sha := range shas {
|
||||
rc, _ := api.Store.DecBlobRef(r.Context(), sha)
|
||||
if rc == 0 {
|
||||
_ = api.Store.DeleteBlob(r.Context(), sha)
|
||||
_ = os.Remove(api.blobFilePath(sha))
|
||||
}
|
||||
}
|
||||
api.Cache.InvalidateRepo(repo.ID)
|
||||
id := middleware.IdentityFromContext(r.Context())
|
||||
_ = api.Store.RecordAudit(r.Context(), &id.User.ID, &repo.ID, "repo.delete", repo.Name, "")
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
// RepoPubkey returns the server's armored public key (convenience endpoint).
|
||||
func (api *API) RepoPubkey(w http.ResponseWriter, r *http.Request) {
|
||||
if api.requireRead(w, r) == nil {
|
||||
return
|
||||
}
|
||||
api.ServerPubkey(w, r)
|
||||
}
|
||||
|
||||
// ListMembers returns the members of a repository.
|
||||
func (api *API) ListMembers(w http.ResponseWriter, r *http.Request) {
|
||||
repo := api.requireRead(w, r)
|
||||
if repo == nil {
|
||||
return
|
||||
}
|
||||
members, err := api.Store.ListMembers(r.Context(), repo.ID)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to list members")
|
||||
return
|
||||
}
|
||||
if members == nil {
|
||||
members = []*models.RepositoryMember{}
|
||||
}
|
||||
httputil.WriteJSON(w, http.StatusOK, members)
|
||||
}
|
||||
|
||||
// AddMember grants a user access on a repository.
|
||||
func (api *API) AddMember(w http.ResponseWriter, r *http.Request) {
|
||||
repo := api.requireManage(w, r)
|
||||
if repo == nil {
|
||||
return
|
||||
}
|
||||
var req apitypes.AddMemberRequest
|
||||
if err := httputil.ReadJSON(r, &req, 1<<20); err != nil {
|
||||
bad(w, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
target, err := api.Store.GetUserByUsername(r.Context(), req.Username)
|
||||
if err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "user not found")
|
||||
return
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to read user")
|
||||
return
|
||||
}
|
||||
if !models.ValidAccess(req.Access) {
|
||||
bad(w, "access must be read, write, read-write, or admin")
|
||||
return
|
||||
}
|
||||
if target.ID == repo.OwnerUserID {
|
||||
bad(w, "cannot change owner's access")
|
||||
return
|
||||
}
|
||||
if err := api.Store.AddMember(r.Context(), repo.ID, target.ID, models.Access(req.Access)); err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to add member")
|
||||
return
|
||||
}
|
||||
_ = api.Store.RecordAudit(r.Context(), &repo.OwnerUserID, &repo.ID, "member.add", req.Username, req.Access)
|
||||
httputil.WriteJSON(w, http.StatusCreated, &models.RepositoryMember{
|
||||
RepositoryID: repo.ID, UserID: target.ID, Access: models.Access(req.Access), User: target,
|
||||
})
|
||||
}
|
||||
|
||||
// UpdateMember changes a member's access level.
|
||||
func (api *API) UpdateMember(w http.ResponseWriter, r *http.Request) {
|
||||
repo := api.requireManage(w, r)
|
||||
if repo == nil {
|
||||
return
|
||||
}
|
||||
var req apitypes.UpdateMemberRequest
|
||||
if err := httputil.ReadJSON(r, &req, 1<<20); err != nil {
|
||||
bad(w, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
if !models.ValidAccess(req.Access) {
|
||||
bad(w, "access must be read, write, read-write, or admin")
|
||||
return
|
||||
}
|
||||
username := r.PathValue("username")
|
||||
target, err := api.Store.GetUserByUsername(r.Context(), username)
|
||||
if err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "user not found")
|
||||
return
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to read user")
|
||||
return
|
||||
}
|
||||
if target.ID == repo.OwnerUserID {
|
||||
bad(w, "cannot change owner's access")
|
||||
return
|
||||
}
|
||||
if err := api.Store.UpdateMemberAccess(r.Context(), repo.ID, target.ID, models.Access(req.Access)); err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "member not found")
|
||||
return
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to update member")
|
||||
return
|
||||
}
|
||||
httputil.WriteJSON(w, http.StatusOK, &models.RepositoryMember{
|
||||
RepositoryID: repo.ID, UserID: target.ID, Access: models.Access(req.Access), User: target,
|
||||
})
|
||||
}
|
||||
|
||||
// RemoveMember revokes a user's access.
|
||||
func (api *API) RemoveMember(w http.ResponseWriter, r *http.Request) {
|
||||
repo := api.requireManage(w, r)
|
||||
if repo == nil {
|
||||
return
|
||||
}
|
||||
username := r.PathValue("username")
|
||||
target, err := api.Store.GetUserByUsername(r.Context(), username)
|
||||
if err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "user not found")
|
||||
return
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to read user")
|
||||
return
|
||||
}
|
||||
if target.ID == repo.OwnerUserID {
|
||||
bad(w, "cannot remove owner")
|
||||
return
|
||||
}
|
||||
if err := api.Store.RemoveMember(r.Context(), repo.ID, target.ID); err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "member not found")
|
||||
return
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to remove member")
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
// valStr returns s as a string pointer-or-nil.
|
||||
func valStr(s *string) string {
|
||||
if s == nil {
|
||||
return ""
|
||||
}
|
||||
return *s
|
||||
}
|
||||
|
||||
// blobFilePath returns the on-disk path for a content-addressed blob.
|
||||
func (api *API) blobFilePath(sha256 string) string {
|
||||
return filepath.Join(api.Config.PackagesDir, sha256+".deb")
|
||||
}
|
||||
@@ -0,0 +1,122 @@
|
||||
package restapi
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRepoCreateAndPermissions(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
// alice (owner/admin), bob (non-admin)
|
||||
_, body := h.do("POST", "/api/v1/auth/register", "", map[string]string{"username": "alice", "password": "supersecret"})
|
||||
var alice struct {
|
||||
User map[string]any `json:"user"`
|
||||
Token string `json:"token"`
|
||||
}
|
||||
json.Unmarshal(body, &alice)
|
||||
|
||||
_, body = h.do("POST", "/api/v1/auth/register", "", map[string]string{"username": "bob", "password": "supersecret"})
|
||||
var bob struct {
|
||||
User map[string]any `json:"user"`
|
||||
Token string `json:"token"`
|
||||
}
|
||||
json.Unmarshal(body, &bob)
|
||||
|
||||
// alice creates a private repo
|
||||
code, body := h.do("POST", "/api/v1/repositories", alice.Token, map[string]any{
|
||||
"name": "myrepo", "visibility": "private", "description": "test",
|
||||
})
|
||||
if code != 201 {
|
||||
t.Fatalf("create repo status %d body %s", code, body)
|
||||
}
|
||||
|
||||
// bob cannot see it (private, not a member)
|
||||
code, body = h.do("GET", "/api/v1/repositories/myrepo", bob.Token, nil)
|
||||
if code != 403 {
|
||||
t.Fatalf("bob should be forbidden from private repo, got %d", code)
|
||||
}
|
||||
|
||||
// alice grants bob read access
|
||||
code, _ = h.do("POST", "/api/v1/repositories/myrepo/members", alice.Token, map[string]string{
|
||||
"username": "bob", "access": "read",
|
||||
})
|
||||
if code != 201 {
|
||||
t.Fatalf("add member status %d", code)
|
||||
}
|
||||
|
||||
// bob can now read
|
||||
code, _ = h.do("GET", "/api/v1/repositories/myrepo", bob.Token, nil)
|
||||
if code != 200 {
|
||||
t.Fatalf("bob should read after grant, got %d", code)
|
||||
}
|
||||
|
||||
// bob cannot write (read only)
|
||||
code, _ = h.do("POST", "/api/v1/repositories/myrepo/distributions", bob.Token, map[string]string{"name": "stable"})
|
||||
if code != 403 {
|
||||
t.Fatalf("bob read-only should not write, got %d", code)
|
||||
}
|
||||
|
||||
// alice upgrades bob to write
|
||||
h.do("PATCH", "/api/v1/repositories/myrepo/members/bob", alice.Token, map[string]string{"access": "write"})
|
||||
code, _ = h.do("POST", "/api/v1/repositories/myrepo/distributions", bob.Token, map[string]string{"name": "stable"})
|
||||
if code != 201 {
|
||||
t.Fatalf("bob with write should create distro, got %d", code)
|
||||
}
|
||||
|
||||
// add component and arch
|
||||
h.do("POST", "/api/v1/repositories/myrepo/distributions/stable/components", alice.Token, map[string]string{"name": "main"})
|
||||
code, _ = h.do("POST", "/api/v1/repositories/myrepo/distributions/stable/architectures", alice.Token, map[string]string{"name": "amd64"})
|
||||
if code != 201 {
|
||||
t.Fatalf("add arch status %d", code)
|
||||
}
|
||||
// 'all' arch rejected
|
||||
code, _ = h.do("POST", "/api/v1/repositories/myrepo/distributions/stable/architectures", alice.Token, map[string]string{"name": "all"})
|
||||
if code != 400 {
|
||||
t.Fatalf("all arch should be rejected, got %d", code)
|
||||
}
|
||||
|
||||
// list distros/components/arches
|
||||
code, _ = h.do("GET", "/api/v1/repositories/myrepo/distributions", alice.Token, nil)
|
||||
if code != 200 {
|
||||
t.Fatalf("list distros status %d", code)
|
||||
}
|
||||
code, _ = h.do("GET", "/api/v1/repositories/myrepo/distributions/stable/components", alice.Token, nil)
|
||||
if code != 200 {
|
||||
t.Fatalf("list components status %d", code)
|
||||
}
|
||||
|
||||
// remove member
|
||||
code, _ = h.do("DELETE", "/api/v1/repositories/myrepo/members/bob", alice.Token, nil)
|
||||
if code != 204 {
|
||||
t.Fatalf("remove member status %d", code)
|
||||
}
|
||||
code, _ = h.do("GET", "/api/v1/repositories/myrepo", bob.Token, nil)
|
||||
if code != 403 {
|
||||
t.Fatalf("bob should be forbidden after removal, got %d", code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublicRepoReadableByAll(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
_, body := h.do("POST", "/api/v1/auth/register", "", map[string]string{"username": "owner", "password": "supersecret"})
|
||||
var owner struct {
|
||||
Token string `json:"token"`
|
||||
}
|
||||
json.Unmarshal(body, &owner)
|
||||
_, body = h.do("POST", "/api/v1/auth/register", "", map[string]string{"username": "stranger", "password": "supersecret"})
|
||||
var stranger struct {
|
||||
Token string `json:"token"`
|
||||
}
|
||||
json.Unmarshal(body, &stranger)
|
||||
|
||||
h.do("POST", "/api/v1/repositories", owner.Token, map[string]any{"name": "pubrepo", "visibility": "public"})
|
||||
code, _ := h.do("GET", "/api/v1/repositories/pubrepo", stranger.Token, nil)
|
||||
if code != 200 {
|
||||
t.Fatalf("stranger should read public repo, got %d", code)
|
||||
}
|
||||
// but stranger cannot write
|
||||
code, _ = h.do("POST", "/api/v1/repositories/pubrepo/distributions", stranger.Token, map[string]string{"name": "x"})
|
||||
if code != 403 {
|
||||
t.Fatalf("stranger should not write public repo, got %d", code)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
package restapi
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
apitypes "urapt/shared/api"
|
||||
"urapt/shared/httputil"
|
||||
"urapt/shared/version"
|
||||
)
|
||||
|
||||
// ServerInfo returns build/setup metadata for the server.
|
||||
func (api *API) ServerInfo(w http.ResponseWriter, r *http.Request) {
|
||||
users, err := api.Store.ListUsers(r.Context())
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to read users")
|
||||
return
|
||||
}
|
||||
fp := ""
|
||||
if api.Signer != nil {
|
||||
fp = api.Signer.Fingerprint()
|
||||
}
|
||||
httputil.WriteJSON(w, http.StatusOK, apitypes.ServerInfo{
|
||||
Version: version.Version,
|
||||
NeedsSetup: len(users) == 0,
|
||||
DefaultKeyFingerprint: fp,
|
||||
OpenRegistration: api.Config.OpenRegistration,
|
||||
})
|
||||
}
|
||||
|
||||
// ServerPubkey returns the ASCII-armored default signing key.
|
||||
func (api *API) ServerPubkey(w http.ResponseWriter, r *http.Request) {
|
||||
if api.Signer == nil {
|
||||
httputil.WriteError(w, http.StatusServiceUnavailable, httputil.CodeInternal, "no signing key configured")
|
||||
return
|
||||
}
|
||||
pub, err := api.Signer.PublicKeyArmored()
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to read key")
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/pgp-keys")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(pub))
|
||||
}
|
||||
@@ -0,0 +1,255 @@
|
||||
package restapi
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"urapt/server/store"
|
||||
apitypes "urapt/shared/api"
|
||||
"urapt/shared/httputil"
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// validateName checks a distribution/component/architecture name.
|
||||
func validateName(s string) bool { return nameRE.MatchString(s) }
|
||||
|
||||
// loadDistro fetches the {repo}/{dist} distribution, rendering errors.
|
||||
func (api *API) loadDistro(w http.ResponseWriter, r *http.Request) (*models.Repository, *models.Distribution) {
|
||||
repo := api.requireRead(w, r)
|
||||
if repo == nil {
|
||||
return nil, nil
|
||||
}
|
||||
dist, err := api.Store.GetDistributionByName(r.Context(), repo.ID, r.PathValue("dist"))
|
||||
if err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "distribution not found")
|
||||
return repo, nil
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to read distribution")
|
||||
return repo, nil
|
||||
}
|
||||
return repo, dist
|
||||
}
|
||||
|
||||
// requireDistroWrite loads repo (write) + distribution.
|
||||
func (api *API) requireDistroWrite(w http.ResponseWriter, r *http.Request) (*models.Repository, *models.Distribution) {
|
||||
repo := api.requireWrite(w, r)
|
||||
if repo == nil {
|
||||
return nil, nil
|
||||
}
|
||||
dist, err := api.Store.GetDistributionByName(r.Context(), repo.ID, r.PathValue("dist"))
|
||||
if err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "distribution not found")
|
||||
return repo, nil
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to read distribution")
|
||||
return repo, nil
|
||||
}
|
||||
return repo, dist
|
||||
}
|
||||
|
||||
// --- distributions ---
|
||||
|
||||
// ListDistributions returns the distributions in a repository.
|
||||
func (api *API) ListDistributions(w http.ResponseWriter, r *http.Request) {
|
||||
repo := api.requireRead(w, r)
|
||||
if repo == nil {
|
||||
return
|
||||
}
|
||||
dists, err := api.Store.ListDistributions(r.Context(), repo.ID)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to list distributions")
|
||||
return
|
||||
}
|
||||
if dists == nil {
|
||||
dists = []*models.Distribution{}
|
||||
}
|
||||
httputil.WriteJSON(w, http.StatusOK, dists)
|
||||
}
|
||||
|
||||
// CreateDistribution adds a distribution to a repository.
|
||||
func (api *API) CreateDistribution(w http.ResponseWriter, r *http.Request) {
|
||||
repo := api.requireWrite(w, r)
|
||||
if repo == nil {
|
||||
return
|
||||
}
|
||||
var req apitypes.CreateNamedRequest
|
||||
if err := httputil.ReadJSON(r, &req, 1<<20); err != nil {
|
||||
bad(w, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
if !validateName(req.Name) {
|
||||
bad(w, "name must be 1-64 chars of [a-z0-9][a-z0-9+.-]")
|
||||
return
|
||||
}
|
||||
if _, err := api.Store.GetDistributionByName(r.Context(), repo.ID, req.Name); err == nil {
|
||||
httputil.WriteError(w, http.StatusConflict, httputil.CodeConflict, "distribution already exists")
|
||||
return
|
||||
} else if !errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to check distribution")
|
||||
return
|
||||
}
|
||||
dist, err := api.Store.CreateDistribution(r.Context(), repo.ID, req.Name)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to create distribution")
|
||||
return
|
||||
}
|
||||
api.Cache.Invalidate(repo.ID, dist.Name)
|
||||
httputil.WriteJSON(w, http.StatusCreated, dist)
|
||||
}
|
||||
|
||||
// DeleteDistribution removes a distribution (cascades to packages).
|
||||
func (api *API) DeleteDistribution(w http.ResponseWriter, r *http.Request) {
|
||||
repo, dist := api.requireDistroWrite(w, r)
|
||||
if dist == nil {
|
||||
return
|
||||
}
|
||||
if err := api.Store.DeleteDistribution(r.Context(), repo.ID, dist.Name); err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "distribution not found")
|
||||
return
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to delete distribution")
|
||||
return
|
||||
}
|
||||
// Note: cascaded package rows are gone; their blob ref counts are now
|
||||
// stale. Best-effort cleanup of orphan blobs is handled by package delete
|
||||
// in normal operation; bulk distribution deletion leaves blobs for now.
|
||||
api.Cache.Invalidate(repo.ID, dist.Name)
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
// --- components ---
|
||||
|
||||
// ListComponents returns the components in a distribution.
|
||||
func (api *API) ListComponents(w http.ResponseWriter, r *http.Request) {
|
||||
_, dist := api.loadDistro(w, r)
|
||||
if dist == nil {
|
||||
return
|
||||
}
|
||||
comps, err := api.Store.ListComponents(r.Context(), dist.ID)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to list components")
|
||||
return
|
||||
}
|
||||
if comps == nil {
|
||||
comps = []*models.Component{}
|
||||
}
|
||||
httputil.WriteJSON(w, http.StatusOK, comps)
|
||||
}
|
||||
|
||||
// CreateComponent adds a component to a distribution.
|
||||
func (api *API) CreateComponent(w http.ResponseWriter, r *http.Request) {
|
||||
repo, dist := api.requireDistroWrite(w, r)
|
||||
if dist == nil {
|
||||
return
|
||||
}
|
||||
var req apitypes.CreateNamedRequest
|
||||
if err := httputil.ReadJSON(r, &req, 1<<20); err != nil {
|
||||
bad(w, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
if !validateName(req.Name) {
|
||||
bad(w, "name must be 1-64 chars of [a-z0-9][a-z0-9+.-]")
|
||||
return
|
||||
}
|
||||
if _, err := api.Store.GetComponentByName(r.Context(), dist.ID, req.Name); err == nil {
|
||||
httputil.WriteError(w, http.StatusConflict, httputil.CodeConflict, "component already exists")
|
||||
return
|
||||
} else if !errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to check component")
|
||||
return
|
||||
}
|
||||
comp, err := api.Store.CreateComponent(r.Context(), dist.ID, req.Name)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to create component")
|
||||
return
|
||||
}
|
||||
api.Cache.Invalidate(repo.ID, dist.Name)
|
||||
httputil.WriteJSON(w, http.StatusCreated, comp)
|
||||
}
|
||||
|
||||
// DeleteComponent removes a component (blocked if packages reference it).
|
||||
func (api *API) DeleteComponent(w http.ResponseWriter, r *http.Request) {
|
||||
repo, dist := api.requireDistroWrite(w, r)
|
||||
if dist == nil {
|
||||
return
|
||||
}
|
||||
comp, err := api.Store.GetComponentByName(r.Context(), dist.ID, r.PathValue("comp"))
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "component not found")
|
||||
return
|
||||
}
|
||||
if err := api.Store.DeleteComponent(r.Context(), dist.ID, comp.Name); err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to delete component (packages may still reference it)")
|
||||
return
|
||||
}
|
||||
api.Cache.Invalidate(repo.ID, dist.Name)
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
// --- architectures ---
|
||||
|
||||
// ListArchitectures returns the architectures in a distribution.
|
||||
func (api *API) ListArchitectures(w http.ResponseWriter, r *http.Request) {
|
||||
_, dist := api.loadDistro(w, r)
|
||||
if dist == nil {
|
||||
return
|
||||
}
|
||||
arches, err := api.Store.ListArchitectures(r.Context(), dist.ID)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to list architectures")
|
||||
return
|
||||
}
|
||||
if arches == nil {
|
||||
arches = []*models.Architecture{}
|
||||
}
|
||||
httputil.WriteJSON(w, http.StatusOK, arches)
|
||||
}
|
||||
|
||||
// CreateArchitecture adds an architecture to a distribution.
|
||||
func (api *API) CreateArchitecture(w http.ResponseWriter, r *http.Request) {
|
||||
repo, dist := api.requireDistroWrite(w, r)
|
||||
if dist == nil {
|
||||
return
|
||||
}
|
||||
var req apitypes.CreateNamedRequest
|
||||
if err := httputil.ReadJSON(r, &req, 1<<20); err != nil {
|
||||
bad(w, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
if !validateName(req.Name) {
|
||||
bad(w, "name must be 1-64 chars of [a-z0-9][a-z0-9+.-]")
|
||||
return
|
||||
}
|
||||
if req.Name == "all" {
|
||||
bad(w, "architecture 'all' is implicit and cannot be added")
|
||||
return
|
||||
}
|
||||
arch, err := api.Store.CreateArchitecture(r.Context(), dist.ID, req.Name)
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to create architecture (already exists?)")
|
||||
return
|
||||
}
|
||||
api.Cache.Invalidate(repo.ID, dist.Name)
|
||||
httputil.WriteJSON(w, http.StatusCreated, arch)
|
||||
}
|
||||
|
||||
// DeleteArchitecture removes an architecture from a distribution.
|
||||
func (api *API) DeleteArchitecture(w http.ResponseWriter, r *http.Request) {
|
||||
repo, dist := api.requireDistroWrite(w, r)
|
||||
if dist == nil {
|
||||
return
|
||||
}
|
||||
if err := api.Store.DeleteArchitecture(r.Context(), dist.ID, r.PathValue("arch")); err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "architecture not found")
|
||||
return
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to delete architecture")
|
||||
return
|
||||
}
|
||||
api.Cache.Invalidate(repo.ID, dist.Name)
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
package restapi
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"urapt/server/middleware"
|
||||
"urapt/server/store"
|
||||
apitypes "urapt/shared/api"
|
||||
"urapt/shared/httputil"
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// ListUsers returns all users (admin only).
|
||||
func (api *API) ListUsers(w http.ResponseWriter, r *http.Request) {
|
||||
users, err := api.Store.ListUsers(r.Context())
|
||||
if err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to list users")
|
||||
return
|
||||
}
|
||||
if users == nil {
|
||||
users = []*models.User{}
|
||||
}
|
||||
httputil.WriteJSON(w, http.StatusOK, apitypes.ListResponse[*models.User]{Items: users, Page: 1, PerPage: 100, Total: len(users)})
|
||||
}
|
||||
|
||||
// GetUser returns a single user by id (admin only).
|
||||
func (api *API) GetUser(w http.ResponseWriter, r *http.Request) {
|
||||
id := r.PathValue("id")
|
||||
user, err := api.Store.GetUserByID(r.Context(), id)
|
||||
if err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "user not found")
|
||||
return
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to read user")
|
||||
return
|
||||
}
|
||||
httputil.WriteJSON(w, http.StatusOK, user)
|
||||
}
|
||||
|
||||
// UpdateUser mutates a user (currently only is_admin) (admin only).
|
||||
func (api *API) UpdateUser(w http.ResponseWriter, r *http.Request) {
|
||||
id := r.PathValue("id")
|
||||
var req apitypes.UpdateUserRequest
|
||||
if err := httputil.ReadJSON(r, &req, 1<<20); err != nil {
|
||||
bad(w, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
target, err := api.Store.GetUserByID(r.Context(), id)
|
||||
if err != nil {
|
||||
if errors.Is(err, store.ErrNotFound) {
|
||||
httputil.WriteError(w, http.StatusNotFound, httputil.CodeNotFound, "user not found")
|
||||
return
|
||||
}
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to read user")
|
||||
return
|
||||
}
|
||||
caller := middleware.IdentityFromContext(r.Context())
|
||||
if req.IsAdmin != nil {
|
||||
if *req.IsAdmin && !caller.User.IsAdmin {
|
||||
httputil.WriteError(w, http.StatusForbidden, httputil.CodeForbidden, "cannot grant admin")
|
||||
return
|
||||
}
|
||||
if target.ID == caller.User.ID && !*req.IsAdmin {
|
||||
httputil.WriteError(w, http.StatusBadRequest, httputil.CodeBadRequest, "cannot revoke your own admin")
|
||||
return
|
||||
}
|
||||
}
|
||||
if err := api.Store.UpdateUser(r.Context(), id, req.IsAdmin); err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to update user")
|
||||
return
|
||||
}
|
||||
updated, _ := api.Store.GetUserByID(r.Context(), id)
|
||||
httputil.WriteJSON(w, http.StatusOK, updated)
|
||||
}
|
||||
|
||||
// DeleteUser removes a user (admin only). Self-deletion is blocked.
|
||||
func (api *API) DeleteUser(w http.ResponseWriter, r *http.Request) {
|
||||
id := r.PathValue("id")
|
||||
caller := middleware.IdentityFromContext(r.Context())
|
||||
if id == caller.User.ID {
|
||||
httputil.WriteError(w, http.StatusBadRequest, httputil.CodeBadRequest, "cannot delete your own account")
|
||||
return
|
||||
}
|
||||
if err := api.Store.DeleteUser(r.Context(), id); err != nil {
|
||||
httputil.WriteError(w, http.StatusInternalServerError, httputil.CodeInternal, "failed to delete user")
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
package store
|
||||
|
||||
import "context"
|
||||
|
||||
// RecordAudit inserts a best-effort audit log entry. Errors are ignored at the
|
||||
// call site's discretion; this helper returns the error for completeness.
|
||||
func (s *Store) RecordAudit(ctx context.Context, userID, repositoryID *string, action, target, details string) error {
|
||||
_, err := s.exec(ctx, `INSERT INTO audit_log (id, user_id, repository_id, action, target, details, created_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)`,
|
||||
newID(), nullable(userID), nullable(repositoryID), action, target, details, s.now())
|
||||
return err
|
||||
}
|
||||
|
||||
func nullable(p *string) any {
|
||||
if p == nil {
|
||||
return nil
|
||||
}
|
||||
return *p
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// GetBlob returns a blob by its sha256, or ErrNotFound.
|
||||
func (s *Store) GetBlob(ctx context.Context, sha256 string) (filename string, size, refCount int64, err error) {
|
||||
err = s.db.QueryRowContext(ctx, `SELECT filename, size, ref_count FROM blobs WHERE sha256 = ?`, sha256).
|
||||
Scan(&filename, &size, &refCount)
|
||||
if isErrNoRows(err) {
|
||||
return "", 0, 0, ErrNotFound
|
||||
}
|
||||
return filename, size, refCount, err
|
||||
}
|
||||
|
||||
// CreateBlob creates a new blob row with ref_count=1. It returns
|
||||
// (created=true) when a new row was inserted, or (created=false) when the
|
||||
// blob already existed (in which case its ref_count is left unchanged here;
|
||||
// use IncBlobRef to bump it).
|
||||
func (s *Store) CreateBlob(ctx context.Context, sha256, filename string, size int64) (created bool, err error) {
|
||||
now := s.now()
|
||||
res, err := s.exec(ctx, `INSERT OR IGNORE INTO blobs (sha256, filename, size, ref_count, created_at) VALUES (?, ?, ?, 1, ?)`,
|
||||
sha256, filename, size, now)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("insert blob: %w", err)
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
return n > 0, nil
|
||||
}
|
||||
|
||||
// IncBlobRef atomically increments a blob's ref_count and returns the new value.
|
||||
func (s *Store) IncBlobRef(ctx context.Context, sha256 string) (int64, error) {
|
||||
res, err := s.exec(ctx, `UPDATE blobs SET ref_count = ref_count + 1 WHERE sha256 = ?`, sha256)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("inc blob: %w", err)
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return 0, ErrNotFound
|
||||
}
|
||||
var rc int64
|
||||
if err := s.db.QueryRowContext(ctx, `SELECT ref_count FROM blobs WHERE sha256 = ?`, sha256).Scan(&rc); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return rc, nil
|
||||
}
|
||||
|
||||
// DecBlobRef atomically decrements a blob's ref_count and returns the new
|
||||
// value. When it reaches 0 the caller should delete the on-disk file and call
|
||||
// DeleteBlob.
|
||||
func (s *Store) DecBlobRef(ctx context.Context, sha256 string) (int64, error) {
|
||||
res, err := s.exec(ctx, `UPDATE blobs SET ref_count = ref_count - 1 WHERE sha256 = ? AND ref_count > 0`, sha256)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("dec blob: %w", err)
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
var rc int64
|
||||
if e := s.db.QueryRowContext(ctx, `SELECT ref_count FROM blobs WHERE sha256 = ?`, sha256).Scan(&rc); e != nil {
|
||||
if isErrNoRows(e) {
|
||||
return 0, ErrNotFound
|
||||
}
|
||||
return 0, e
|
||||
}
|
||||
return rc, nil
|
||||
}
|
||||
var rc int64
|
||||
if err := s.db.QueryRowContext(ctx, `SELECT ref_count FROM blobs WHERE sha256 = ?`, sha256).Scan(&rc); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return rc, nil
|
||||
}
|
||||
|
||||
// DeleteBlob removes a blob row.
|
||||
func (s *Store) DeleteBlob(ctx context.Context, sha256 string) error {
|
||||
_, err := s.exec(ctx, `DELETE FROM blobs WHERE sha256 = ?`, sha256)
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// SaveGPGKey inserts a GPG key row.
|
||||
func (s *Store) SaveGPGKey(ctx context.Context, fingerprint, userID, pubArmored, privArmored string, isDefault bool) (*models.GPGKey, error) {
|
||||
id := newID()
|
||||
now := s.now()
|
||||
_, err := s.exec(ctx, `INSERT INTO gpg_keys (id, fingerprint, user_id, public_key_armored, private_key_armored, is_default, created_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)`,
|
||||
id, fingerprint, userID, pubArmored, privArmored, boolToInt(isDefault), now)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("insert gpg key: %w", err)
|
||||
}
|
||||
return &models.GPGKey{
|
||||
ID: id, Fingerprint: fingerprint, UserID: userID,
|
||||
PublicKeyArmored: pubArmored, IsDefault: isDefault, CreatedAt: now,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetDefaultGPGKey returns the default signing key, or ErrNotFound if none.
|
||||
func (s *Store) GetDefaultGPGKey(ctx context.Context) (*models.GPGKey, error) {
|
||||
row := s.db.QueryRowContext(ctx,
|
||||
`SELECT id, fingerprint, user_id, public_key_armored, private_key_armored, is_default, created_at
|
||||
FROM gpg_keys WHERE is_default = 1 LIMIT 1`)
|
||||
k := &models.GPGKey{}
|
||||
var isDefault int
|
||||
err := row.Scan(&k.ID, &k.Fingerprint, &k.UserID, &k.PublicKeyArmored, &k.PrivateKeyArmored, &isDefault, &k.CreatedAt)
|
||||
if isErrNoRows(err) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
k.IsDefault = isDefault == 1
|
||||
return k, nil
|
||||
}
|
||||
|
||||
// GetGPGKeyPublic returns just the armored public key of the default key.
|
||||
func (s *Store) GetGPGKeyPublic(ctx context.Context) (string, error) {
|
||||
row := s.db.QueryRowContext(ctx, `SELECT public_key_armored FROM gpg_keys WHERE is_default = 1 LIMIT 1`)
|
||||
var pub string
|
||||
err := row.Scan(&pub)
|
||||
if isErrNoRows(err) {
|
||||
return "", ErrNotFound
|
||||
}
|
||||
return pub, err
|
||||
}
|
||||
@@ -0,0 +1,102 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSaveGPGKeyAndGetDefault(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
|
||||
k, err := s.SaveGPGKey(ctx, "ABCD1234", "user-1", "PUB-ARMORED", "PRIV-ARMORED", true)
|
||||
if err != nil {
|
||||
t.Fatalf("save: %v", err)
|
||||
}
|
||||
if k.ID == "" || k.Fingerprint != "ABCD1234" || !k.IsDefault {
|
||||
t.Fatalf("key = %+v", k)
|
||||
}
|
||||
|
||||
got, err := s.GetDefaultGPGKey(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("get default: %v", err)
|
||||
}
|
||||
if got.Fingerprint != "ABCD1234" || got.PublicKeyArmored != "PUB-ARMORED" {
|
||||
t.Fatalf("got = %+v", got)
|
||||
}
|
||||
if got.PrivateKeyArmored != "PRIV-ARMORED" {
|
||||
t.Fatal("private key armor should be retrieved from DB")
|
||||
}
|
||||
if !got.IsDefault {
|
||||
t.Fatal("expected is_default=true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetDefaultGPGKey_None(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
if _, err := s.GetDefaultGPGKey(ctx); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound when no key, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetGPGKeyPublic(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
_, _ = s.SaveGPGKey(ctx, "ABCD1234", "user-1", "PUB-ARMORED", "PRIV-ARMORED", true)
|
||||
|
||||
pub, err := s.GetGPGKeyPublic(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("get public: %v", err)
|
||||
}
|
||||
if pub != "PUB-ARMORED" {
|
||||
t.Fatalf("pub = %q", pub)
|
||||
}
|
||||
// No default key.
|
||||
s2 := newTestStore(t)
|
||||
if _, err := s2.GetGPGKeyPublic(ctx); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// --- audit ---
|
||||
|
||||
func TestRecordAudit(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
// audit_log has FK constraints on user_id/repo_id, so use real rows.
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, "public")
|
||||
uid, repoID := u.ID, repo.ID
|
||||
|
||||
// Insert with both user and repo set.
|
||||
if err := s.RecordAudit(ctx, &uid, &repoID, "push", "pkg-1", "uploaded myapp"); err != nil {
|
||||
t.Fatalf("record: %v", err)
|
||||
}
|
||||
// Insert with nil pointers (system action).
|
||||
if err := s.RecordAudit(ctx, nil, nil, "startup", "server", "started"); err != nil {
|
||||
t.Fatalf("record nil: %v", err)
|
||||
}
|
||||
|
||||
// Verify rows exist. The audit_log table is write-only from the store API;
|
||||
// query directly to confirm persistence.
|
||||
var n int
|
||||
err := s.DB().QueryRowContext(ctx, `SELECT COUNT(*) FROM audit_log`).Scan(&n)
|
||||
if err != nil {
|
||||
t.Fatalf("count: %v", err)
|
||||
}
|
||||
if n != 2 {
|
||||
t.Fatalf("expected 2 audit rows, got %d", n)
|
||||
}
|
||||
|
||||
// Verify nullable columns are stored correctly.
|
||||
var uidVal, repoVal *string
|
||||
row := s.DB().QueryRowContext(ctx, `SELECT user_id, repository_id FROM audit_log WHERE action = 'startup'`)
|
||||
if err := row.Scan(&uidVal, &repoVal); err != nil {
|
||||
t.Fatalf("scan startup row: %v", err)
|
||||
}
|
||||
if uidVal != nil || repoVal != nil {
|
||||
t.Fatalf("expected nil user_id/repository_id for system action, got %v %v", uidVal, repoVal)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// AddMember grants a user access on a repository. If a grant already exists it
|
||||
// is updated.
|
||||
func (s *Store) AddMember(ctx context.Context, repoID, userID string, access models.Access) error {
|
||||
now := s.now()
|
||||
_, err := s.exec(ctx, `INSERT INTO repository_members (repository_id, user_id, access, created_at)
|
||||
VALUES (?, ?, ?, ?)
|
||||
ON CONFLICT(repository_id, user_id) DO UPDATE SET access = excluded.access`,
|
||||
repoID, userID, string(access), now)
|
||||
if err != nil {
|
||||
return fmt.Errorf("upsert member: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateMemberAccess changes a user's access on a repository.
|
||||
func (s *Store) UpdateMemberAccess(ctx context.Context, repoID, userID string, access models.Access) error {
|
||||
res, err := s.exec(ctx, `UPDATE repository_members SET access = ? WHERE repository_id = ? AND user_id = ?`,
|
||||
string(access), repoID, userID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("update member: %w", err)
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return ErrNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveMember revokes a user's access on a repository.
|
||||
func (s *Store) RemoveMember(ctx context.Context, repoID, userID string) error {
|
||||
res, err := s.exec(ctx, `DELETE FROM repository_members WHERE repository_id = ? AND user_id = ?`, repoID, userID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("remove member: %w", err)
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return ErrNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListMembers returns all members of a repository with their user records.
|
||||
func (s *Store) ListMembers(ctx context.Context, repoID string) ([]*models.RepositoryMember, error) {
|
||||
rows, err := s.db.QueryContext(ctx, `
|
||||
SELECT m.repository_id, m.user_id, m.access, m.created_at,
|
||||
u.id, u.username, u.is_admin, u.created_at, u.updated_at
|
||||
FROM repository_members m
|
||||
JOIN users u ON u.id = m.user_id
|
||||
WHERE m.repository_id = ?
|
||||
ORDER BY u.username`, repoID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list members: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []*models.RepositoryMember
|
||||
for rows.Next() {
|
||||
m := &models.RepositoryMember{User: &models.User{}}
|
||||
if err := rows.Scan(
|
||||
&m.RepositoryID, &m.UserID, &m.Access, &m.CreatedAt,
|
||||
&m.User.ID, &m.User.Username, &m.User.IsAdmin, &m.User.CreatedAt, &m.User.UpdatedAt,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, m)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
@@ -0,0 +1,236 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// SuitePackage is a package row joined with its component and distribution
|
||||
// names, used by the APT index generator.
|
||||
type SuitePackage struct {
|
||||
models.Package
|
||||
ComponentName string
|
||||
DistributionName string
|
||||
}
|
||||
|
||||
// CreatePackage inserts a new package row.
|
||||
func (s *Store) CreatePackage(ctx context.Context, p *models.Package) error {
|
||||
if p.ID == "" {
|
||||
p.ID = newID()
|
||||
}
|
||||
if p.CreatedAt == "" {
|
||||
p.CreatedAt = s.now()
|
||||
}
|
||||
_, err := s.exec(ctx, `INSERT INTO packages (
|
||||
id, repository_id, distribution_id, component_id,
|
||||
name, version, architecture, source, maintainer, priority, section,
|
||||
origin, homepage, description, description_md5, depends, pre_depends,
|
||||
recommends, suggests, conflicts, breaks, provides, replaces, enhances,
|
||||
installed_size, essential, built_using, tag, raw_control, filename,
|
||||
pool_path, size, md5sum, sha1, sha256, uploaded_by_user_id, created_at
|
||||
) VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)`,
|
||||
p.ID, p.RepositoryID, p.DistributionID, p.ComponentID,
|
||||
p.Name, p.Version, p.Architecture, p.Source, p.Maintainer, p.Priority, p.Section,
|
||||
p.Origin, p.Homepage, p.Description, p.DescriptionMD5, p.Depends, p.PreDepends,
|
||||
p.Recommends, p.Suggests, p.Conflicts, p.Breaks, p.Provides, p.Replaces, p.Enhances,
|
||||
p.InstalledSize, p.Essential, p.BuiltUsing, p.Tag, p.RawControl, p.Filename,
|
||||
p.PoolPath, p.Size, p.MD5sum, p.SHA1, p.SHA256, p.UploadedByUserID, p.CreatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("insert package: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetPackageByID returns a package by id.
|
||||
func (s *Store) GetPackageByID(ctx context.Context, id string) (*models.Package, error) {
|
||||
row := s.db.QueryRowContext(ctx, packageCols+` FROM packages WHERE id = ?`, id)
|
||||
return scanPackage(row)
|
||||
}
|
||||
|
||||
// GetPackageByPoolPath returns a package within a repository by its pool_path.
|
||||
func (s *Store) GetPackageByPoolPath(ctx context.Context, repoID, poolPath string) (*models.Package, error) {
|
||||
row := s.db.QueryRowContext(ctx, packageCols+` FROM packages WHERE repository_id = ? AND pool_path = ?`, repoID, poolPath)
|
||||
return scanPackage(row)
|
||||
}
|
||||
|
||||
// PackageFilters controls ListPackages filtering.
|
||||
type PackageFilters struct {
|
||||
ComponentID string
|
||||
Arch string
|
||||
Name string
|
||||
Query string
|
||||
}
|
||||
|
||||
// ListPackages lists packages within a (repo, distribution) with optional
|
||||
// filters and pagination.
|
||||
func (s *Store) ListPackages(ctx context.Context, repoID, distroID string, f PackageFilters, page, perPage int) ([]*models.Package, int, error) {
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if perPage < 1 || perPage > 100 {
|
||||
perPage = 25
|
||||
}
|
||||
var where []string
|
||||
var args []any
|
||||
where = append(where, "repository_id = ?", "distribution_id = ?")
|
||||
args = append(args, repoID, distroID)
|
||||
if f.ComponentID != "" {
|
||||
where = append(where, "component_id = ?")
|
||||
args = append(args, f.ComponentID)
|
||||
}
|
||||
if f.Arch != "" {
|
||||
where = append(where, "(architecture = ? OR architecture = 'all')")
|
||||
args = append(args, f.Arch)
|
||||
}
|
||||
if f.Name != "" {
|
||||
where = append(where, "name = ?")
|
||||
args = append(args, f.Name)
|
||||
}
|
||||
if f.Query != "" {
|
||||
where = append(where, "(name LIKE ? OR description LIKE ?)")
|
||||
args = append(args, "%"+f.Query+"%", "%"+f.Query+"%")
|
||||
}
|
||||
q := strings.Join(where, " AND ")
|
||||
|
||||
var total int
|
||||
if err := s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM packages WHERE `+q, args...).Scan(&total); err != nil {
|
||||
return nil, 0, fmt.Errorf("count packages: %w", err)
|
||||
}
|
||||
args2 := append(args, perPage, (page-1)*perPage)
|
||||
rows, err := s.db.QueryContext(ctx, packageCols+` FROM packages WHERE `+q+` ORDER BY name, version LIMIT ? OFFSET ?`, args2...)
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("list packages: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []*models.Package
|
||||
for rows.Next() {
|
||||
p, err := scanPackageRows(rows)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
out = append(out, p)
|
||||
}
|
||||
return out, total, rows.Err()
|
||||
}
|
||||
|
||||
// ListSuitePackages returns all packages in a (repo, distribution) joined with
|
||||
// their component and distribution names, for APT index generation.
|
||||
func (s *Store) ListSuitePackages(ctx context.Context, repoID, distroID string) ([]SuitePackage, error) {
|
||||
cols := qualifiedPackageCols("p")
|
||||
rows, err := s.db.QueryContext(ctx, `
|
||||
SELECT `+cols+`, c.name, d.name
|
||||
FROM packages p
|
||||
JOIN components c ON c.id = p.component_id
|
||||
JOIN distributions d ON d.id = p.distribution_id
|
||||
WHERE p.repository_id = ? AND p.distribution_id = ?`, repoID, distroID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list suite packages: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []SuitePackage
|
||||
for rows.Next() {
|
||||
var sp SuitePackage
|
||||
if err := scanPackageColsFull(rows.Scan, &sp.Package, &sp.ComponentName, &sp.DistributionName); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, sp)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// scanPackageColsFull scans the 37 package columns plus the joined component
|
||||
// name and distribution name.
|
||||
func scanPackageColsFull(scan scanFunc, p *models.Package, compName, distName *string) error {
|
||||
return scan(
|
||||
&p.ID, &p.RepositoryID, &p.DistributionID, &p.ComponentID,
|
||||
&p.Name, &p.Version, &p.Architecture, &p.Source, &p.Maintainer, &p.Priority, &p.Section,
|
||||
&p.Origin, &p.Homepage, &p.Description, &p.DescriptionMD5, &p.Depends, &p.PreDepends,
|
||||
&p.Recommends, &p.Suggests, &p.Conflicts, &p.Breaks, &p.Provides, &p.Replaces, &p.Enhances,
|
||||
&p.InstalledSize, &p.Essential, &p.BuiltUsing, &p.Tag, &p.RawControl, &p.Filename,
|
||||
&p.PoolPath, &p.Size, &p.MD5sum, &p.SHA1, &p.SHA256, &p.UploadedByUserID, &p.CreatedAt,
|
||||
compName, distName,
|
||||
)
|
||||
}
|
||||
|
||||
// ListPackageBlobSHA256sByRepo returns the sha256 of every package in a repo
|
||||
// (with duplicates), used during repository deletion to decrement ref counts.
|
||||
func (s *Store) ListPackageBlobSHA256sByRepo(ctx context.Context, repoID string) ([]string, error) {
|
||||
rows, err := s.db.QueryContext(ctx, `SELECT sha256 FROM packages WHERE repository_id = ?`, repoID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list repo blobs: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []string
|
||||
for rows.Next() {
|
||||
var sha string
|
||||
if err := rows.Scan(&sha); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, sha)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// DeletePackage removes a package row by id.
|
||||
func (s *Store) DeletePackage(ctx context.Context, id string) error {
|
||||
_, err := s.exec(ctx, `DELETE FROM packages WHERE id = ?`, id)
|
||||
return err
|
||||
}
|
||||
|
||||
// DeletePackagesByRepo removes all package rows in a repository.
|
||||
func (s *Store) DeletePackagesByRepo(ctx context.Context, repoID string) error {
|
||||
_, err := s.exec(ctx, `DELETE FROM packages WHERE repository_id = ?`, repoID)
|
||||
return err
|
||||
}
|
||||
|
||||
const packageColNames = `id, repository_id, distribution_id, component_id, name, version, architecture, source, maintainer, priority, section, origin, homepage, description, description_md5, depends, pre_depends, recommends, suggests, conflicts, breaks, provides, replaces, enhances, installed_size, essential, built_using, tag, raw_control, filename, pool_path, size, md5sum, sha1, sha256, uploaded_by_user_id, created_at`
|
||||
|
||||
const packageCols = `SELECT ` + packageColNames
|
||||
|
||||
// qualifiedPackageCols returns the package column list prefixed with alias.,
|
||||
// e.g. "p.id, p.repository_id, ...", for use in JOINs.
|
||||
func qualifiedPackageCols(alias string) string {
|
||||
parts := strings.Split(packageColNames, ", ")
|
||||
for i, p := range parts {
|
||||
parts[i] = alias + "." + p
|
||||
}
|
||||
return strings.Join(parts, ", ")
|
||||
}
|
||||
|
||||
func scanPackage(row *sql.Row) (*models.Package, error) {
|
||||
p := &models.Package{}
|
||||
if err := scanPackageCols(row.Scan, p); err != nil {
|
||||
if isErrNoRows(err) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return p, nil
|
||||
}
|
||||
|
||||
func scanPackageRows(rows *sql.Rows) (*models.Package, error) {
|
||||
p := &models.Package{}
|
||||
if err := scanPackageCols(rows.Scan, p); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return p, nil
|
||||
}
|
||||
|
||||
// scanFunc abstracts *sql.Row.Scan and *sql.Rows.Scan.
|
||||
type scanFunc func(dest ...any) error
|
||||
|
||||
func scanPackageCols(scan scanFunc, p *models.Package) error {
|
||||
return scan(
|
||||
&p.ID, &p.RepositoryID, &p.DistributionID, &p.ComponentID,
|
||||
&p.Name, &p.Version, &p.Architecture, &p.Source, &p.Maintainer, &p.Priority, &p.Section,
|
||||
&p.Origin, &p.Homepage, &p.Description, &p.DescriptionMD5, &p.Depends, &p.PreDepends,
|
||||
&p.Recommends, &p.Suggests, &p.Conflicts, &p.Breaks, &p.Provides, &p.Replaces, &p.Enhances,
|
||||
&p.InstalledSize, &p.Essential, &p.BuiltUsing, &p.Tag, &p.RawControl, &p.Filename,
|
||||
&p.PoolPath, &p.Size, &p.MD5sum, &p.SHA1, &p.SHA256, &p.UploadedByUserID, &p.CreatedAt,
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,366 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// newPackage creates a minimal valid Package row ready to insert.
|
||||
func newPackage(repoID, distroID, compID, userID string) *models.Package {
|
||||
return &models.Package{
|
||||
RepositoryID: repoID,
|
||||
DistributionID: distroID,
|
||||
ComponentID: compID,
|
||||
Name: "myapp-hello",
|
||||
Version: "1.0.0",
|
||||
Architecture: "amd64",
|
||||
Maintainer: "Test <test@example.com>",
|
||||
Description: "a test package",
|
||||
RawControl: "Package: myapp-hello\nVersion: 1.0.0\n",
|
||||
Filename: "myapp-hello_1.0.0_amd64.deb",
|
||||
PoolPath: "pool/main/m/myapp-hello/myapp-hello_1.0.0_amd64.deb",
|
||||
Size: 712,
|
||||
MD5sum: "d41d8cd98f00b204e9800998ecf8427e",
|
||||
SHA1: "da39a3ee5e6b4b0d3255bfef95601890afd80709",
|
||||
SHA256: "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855",
|
||||
UploadedByUserID: userID,
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreatePackageAndGetByID(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, models.VisibilityPublic)
|
||||
d := s.createDistro(t, ctx, repo.ID, "stable")
|
||||
c := s.createComponent(t, ctx, d.ID, "main")
|
||||
_ = c
|
||||
|
||||
p := newPackage(repo.ID, d.ID, c.ID, u.ID)
|
||||
if err := s.CreatePackage(ctx, p); err != nil {
|
||||
t.Fatalf("create: %v", err)
|
||||
}
|
||||
if p.ID == "" || p.CreatedAt == "" {
|
||||
t.Fatal("id/created_at should be populated by CreatePackage")
|
||||
}
|
||||
|
||||
got, err := s.GetPackageByID(ctx, p.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("get by id: %v", err)
|
||||
}
|
||||
if got.Name != "myapp-hello" || got.Version != "1.0.0" || got.Architecture != "amd64" {
|
||||
t.Fatalf("got = %+v", got)
|
||||
}
|
||||
if got.SHA256 != p.SHA256 {
|
||||
t.Fatalf("sha256 mismatch")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetPackageByID_NotFound(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
if _, err := s.GetPackageByID(ctx, "nope"); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetPackageByPoolPath(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, models.VisibilityPublic)
|
||||
d := s.createDistro(t, ctx, repo.ID, "stable")
|
||||
c := s.createComponent(t, ctx, d.ID, "main")
|
||||
|
||||
p := newPackage(repo.ID, d.ID, c.ID, u.ID)
|
||||
_ = s.CreatePackage(ctx, p)
|
||||
|
||||
got, err := s.GetPackageByPoolPath(ctx, repo.ID, p.PoolPath)
|
||||
if err != nil {
|
||||
t.Fatalf("get by pool path: %v", err)
|
||||
}
|
||||
if got.ID != p.ID {
|
||||
t.Fatal("id mismatch")
|
||||
}
|
||||
// Wrong repo.
|
||||
if _, err := s.GetPackageByPoolPath(ctx, "other-repo", p.PoolPath); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound for wrong repo, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListPackages_FiltersAndPagination(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, models.VisibilityPublic)
|
||||
d := s.createDistro(t, ctx, repo.ID, "stable")
|
||||
main := s.createComponent(t, ctx, d.ID, "main")
|
||||
contrib := s.createComponent(t, ctx, d.ID, "contrib")
|
||||
|
||||
// Insert 5 packages: 3 in main, 2 in contrib; mix of amd64 and arm64.
|
||||
pkgs := []*models.Package{
|
||||
newPackage(repo.ID, d.ID, main.ID, u.ID),
|
||||
newPackage(repo.ID, d.ID, main.ID, u.ID),
|
||||
newPackage(repo.ID, d.ID, main.ID, u.ID),
|
||||
newPackage(repo.ID, d.ID, contrib.ID, u.ID),
|
||||
newPackage(repo.ID, d.ID, contrib.ID, u.ID),
|
||||
}
|
||||
names := []string{"alpha", "beta", "gamma", "delta", "epsilon"}
|
||||
arches := []string{"amd64", "amd64", "arm64", "amd64", "arm64"}
|
||||
for i, p := range pkgs {
|
||||
p.Name = names[i]
|
||||
p.Architecture = arches[i]
|
||||
p.PoolPath = "pool/main/m/" + names[i] + "/" + names[i] + "_1.0.0_" + arches[i] + ".deb"
|
||||
p.Filename = names[i] + "_1.0.0_" + arches[i] + ".deb"
|
||||
if err := s.CreatePackage(ctx, p); err != nil {
|
||||
t.Fatalf("create %d: %v", i, err)
|
||||
}
|
||||
}
|
||||
|
||||
// All.
|
||||
list, total, err := s.ListPackages(ctx, repo.ID, d.ID, PackageFilters{}, 1, 100)
|
||||
if err != nil {
|
||||
t.Fatalf("list all: %v", err)
|
||||
}
|
||||
if total != 5 || len(list) != 5 {
|
||||
t.Fatalf("expected total=5 len=5, got total=%d len=%d", total, len(list))
|
||||
}
|
||||
|
||||
// Filter by component.
|
||||
list, total, _ = s.ListPackages(ctx, repo.ID, d.ID, PackageFilters{ComponentID: main.ID}, 1, 100)
|
||||
if total != 3 || len(list) != 3 {
|
||||
t.Fatalf("main filter: expected 3, got total=%d len=%d", total, len(list))
|
||||
}
|
||||
|
||||
// Filter by arch (matches arch OR 'all').
|
||||
list, total, _ = s.ListPackages(ctx, repo.ID, d.ID, PackageFilters{Arch: "amd64"}, 1, 100)
|
||||
if total != 3 { // alpha, beta, delta
|
||||
t.Fatalf("amd64 filter: expected 3, got %d", total)
|
||||
}
|
||||
list, total, _ = s.ListPackages(ctx, repo.ID, d.ID, PackageFilters{Arch: "arm64"}, 1, 100)
|
||||
if total != 2 { // gamma, epsilon
|
||||
t.Fatalf("arm64 filter: expected 2, got %d", total)
|
||||
}
|
||||
|
||||
// Filter by exact name.
|
||||
list, total, _ = s.ListPackages(ctx, repo.ID, d.ID, PackageFilters{Name: "beta"}, 1, 100)
|
||||
if total != 1 || len(list) != 1 || list[0].Name != "beta" {
|
||||
t.Fatalf("name filter: total=%d list=%v", total, list)
|
||||
}
|
||||
|
||||
// Query (LIKE on name/description).
|
||||
list, total, _ = s.ListPackages(ctx, repo.ID, d.ID, PackageFilters{Query: "a test"}, 1, 100)
|
||||
if total != 5 {
|
||||
t.Fatalf("query filter: expected all 5 to match description, got %d", total)
|
||||
}
|
||||
|
||||
// Pagination: page 1, perPage 2 -> 2 items, total 5.
|
||||
list, total, _ = s.ListPackages(ctx, repo.ID, d.ID, PackageFilters{}, 1, 2)
|
||||
if total != 5 || len(list) != 2 {
|
||||
t.Fatalf("page 1: total=%d len=%d", total, len(list))
|
||||
}
|
||||
// page 3 -> only 1 item.
|
||||
list, total, _ = s.ListPackages(ctx, repo.ID, d.ID, PackageFilters{}, 3, 2)
|
||||
if len(list) != 1 {
|
||||
t.Fatalf("page 3: expected 1 item, got %d", len(list))
|
||||
}
|
||||
_ = list
|
||||
}
|
||||
|
||||
func TestDeletePackage(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, models.VisibilityPublic)
|
||||
d := s.createDistro(t, ctx, repo.ID, "stable")
|
||||
c := s.createComponent(t, ctx, d.ID, "main")
|
||||
p := newPackage(repo.ID, d.ID, c.ID, u.ID)
|
||||
_ = s.CreatePackage(ctx, p)
|
||||
|
||||
if err := s.DeletePackage(ctx, p.ID); err != nil {
|
||||
t.Fatalf("delete: %v", err)
|
||||
}
|
||||
if _, err := s.GetPackageByID(ctx, p.ID); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound after delete, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeletePackagesByRepo(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, models.VisibilityPublic)
|
||||
d := s.createDistro(t, ctx, repo.ID, "stable")
|
||||
c := s.createComponent(t, ctx, d.ID, "main")
|
||||
_ = s.CreatePackage(ctx, newPackage(repo.ID, d.ID, c.ID, u.ID))
|
||||
_ = s.CreatePackage(ctx, newPackage(repo.ID, d.ID, c.ID, u.ID))
|
||||
|
||||
if err := s.DeletePackagesByRepo(ctx, repo.ID); err != nil {
|
||||
t.Fatalf("delete by repo: %v", err)
|
||||
}
|
||||
list, total, _ := s.ListPackages(ctx, repo.ID, d.ID, PackageFilters{}, 1, 100)
|
||||
if total != 0 || len(list) != 0 {
|
||||
t.Fatalf("expected no packages after DeletePackagesByRepo, got total=%d", total)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListSuitePackages(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, models.VisibilityPublic)
|
||||
d := s.createDistro(t, ctx, repo.ID, "stable")
|
||||
c := s.createComponent(t, ctx, d.ID, "main")
|
||||
p := newPackage(repo.ID, d.ID, c.ID, u.ID)
|
||||
_ = s.CreatePackage(ctx, p)
|
||||
|
||||
suite, err := s.ListSuitePackages(ctx, repo.ID, d.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("list suite: %v", err)
|
||||
}
|
||||
if len(suite) != 1 {
|
||||
t.Fatalf("expected 1 suite package, got %d", len(suite))
|
||||
}
|
||||
if suite[0].ComponentName != "main" || suite[0].DistributionName != "stable" {
|
||||
t.Fatalf("component=%q distro=%q", suite[0].ComponentName, suite[0].DistributionName)
|
||||
}
|
||||
if suite[0].Name != "myapp-hello" {
|
||||
t.Fatalf("name=%q", suite[0].Name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListPackageBlobSHA256sByRepo(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, models.VisibilityPublic)
|
||||
d := s.createDistro(t, ctx, repo.ID, "stable")
|
||||
c := s.createComponent(t, ctx, d.ID, "main")
|
||||
|
||||
p1 := newPackage(repo.ID, d.ID, c.ID, u.ID)
|
||||
p1.Name = "alpha"
|
||||
p1.SHA256 = "aaa"
|
||||
p1.PoolPath = "pool/main/m/alpha/alpha.deb"
|
||||
p1.Filename = "alpha.deb"
|
||||
_ = s.CreatePackage(ctx, p1)
|
||||
p2 := newPackage(repo.ID, d.ID, c.ID, u.ID)
|
||||
p2.Name = "beta"
|
||||
p2.SHA256 = "bbb"
|
||||
p2.PoolPath = "pool/main/m/beta/beta.deb"
|
||||
p2.Filename = "beta.deb"
|
||||
_ = s.CreatePackage(ctx, p2)
|
||||
|
||||
shas, err := s.ListPackageBlobSHA256sByRepo(ctx, repo.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("list: %v", err)
|
||||
}
|
||||
if len(shas) != 2 {
|
||||
t.Fatalf("expected 2 shas, got %d", len(shas))
|
||||
}
|
||||
}
|
||||
|
||||
// --- blobs ---
|
||||
|
||||
func TestCreateBlob_NewAndExisting(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
|
||||
created, err := s.CreateBlob(ctx, "sha-aaa", "hello.deb", 712)
|
||||
if err != nil {
|
||||
t.Fatalf("create: %v", err)
|
||||
}
|
||||
if !created {
|
||||
t.Fatal("first CreateBlob should report created=true")
|
||||
}
|
||||
// Second time, same sha -> not created, no error.
|
||||
created, err = s.CreateBlob(ctx, "sha-aaa", "hello.deb", 712)
|
||||
if err != nil {
|
||||
t.Fatalf("second create: %v", err)
|
||||
}
|
||||
if created {
|
||||
t.Fatal("second CreateBlob should report created=false (already exists)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetBlob(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
_, _ = s.CreateBlob(ctx, "sha-aaa", "hello.deb", 712)
|
||||
|
||||
filename, size, refCount, err := s.GetBlob(ctx, "sha-aaa")
|
||||
if err != nil {
|
||||
t.Fatalf("get: %v", err)
|
||||
}
|
||||
if filename != "hello.deb" || size != 712 || refCount != 1 {
|
||||
t.Fatalf("got filename=%q size=%d ref=%d", filename, size, refCount)
|
||||
}
|
||||
if _, _, _, err := s.GetBlob(ctx, "missing"); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIncBlobRef(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
_, _ = s.CreateBlob(ctx, "sha-aaa", "hello.deb", 712)
|
||||
|
||||
rc, err := s.IncBlobRef(ctx, "sha-aaa")
|
||||
if err != nil {
|
||||
t.Fatalf("inc: %v", err)
|
||||
}
|
||||
if rc != 2 {
|
||||
t.Fatalf("expected ref=2 after inc, got %d", rc)
|
||||
}
|
||||
rc, _ = s.IncBlobRef(ctx, "sha-aaa")
|
||||
if rc != 3 {
|
||||
t.Fatalf("expected ref=3, got %d", rc)
|
||||
}
|
||||
// Inc on missing blob.
|
||||
if _, err := s.IncBlobRef(ctx, "missing"); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecBlobRef(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
_, _ = s.CreateBlob(ctx, "sha-aaa", "hello.deb", 712)
|
||||
_, _ = s.IncBlobRef(ctx, "sha-aaa") // ref=2
|
||||
|
||||
rc, err := s.DecBlobRef(ctx, "sha-aaa")
|
||||
if err != nil {
|
||||
t.Fatalf("dec: %v", err)
|
||||
}
|
||||
if rc != 1 {
|
||||
t.Fatalf("expected ref=1 after dec, got %d", rc)
|
||||
}
|
||||
rc, _ = s.DecBlobRef(ctx, "sha-aaa")
|
||||
if rc != 0 {
|
||||
t.Fatalf("expected ref=0, got %d", rc)
|
||||
}
|
||||
// Dec below 0 should not go negative; ref_count > 0 guard.
|
||||
rc, _ = s.DecBlobRef(ctx, "sha-aaa")
|
||||
if rc != 0 {
|
||||
t.Fatalf("expected ref to stay at 0, got %d", rc)
|
||||
}
|
||||
// Dec on missing blob.
|
||||
if _, err := s.DecBlobRef(ctx, "totally-missing"); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound for missing, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteBlob(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
_, _ = s.CreateBlob(ctx, "sha-aaa", "hello.deb", 712)
|
||||
|
||||
if err := s.DeleteBlob(ctx, "sha-aaa"); err != nil {
|
||||
t.Fatalf("delete: %v", err)
|
||||
}
|
||||
if _, _, _, err := s.GetBlob(ctx, "sha-aaa"); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound after delete, got %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// CreateRepository inserts a new repository owned by userID.
|
||||
func (s *Store) CreateRepository(ctx context.Context, name, ownerUserID string, visibility models.Visibility, description string) (*models.Repository, error) {
|
||||
id := newID()
|
||||
now := s.now()
|
||||
_, err := s.exec(ctx, `INSERT INTO repositories (id, name, owner_user_id, visibility, description, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)`, id, name, ownerUserID, string(visibility), description, now, now)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("insert repository: %w", err)
|
||||
}
|
||||
return &models.Repository{
|
||||
ID: id, Name: name, OwnerUserID: ownerUserID, Visibility: visibility,
|
||||
Description: description, CreatedAt: now, UpdatedAt: now,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ListReposVisible returns repositories visible to userID: owned, public, or
|
||||
// where the user is a member.
|
||||
func (s *Store) ListReposVisible(ctx context.Context, userID string) ([]*models.Repository, error) {
|
||||
rows, err := s.db.QueryContext(ctx, `
|
||||
SELECT DISTINCT r.id, r.name, r.owner_user_id, r.visibility, r.description, r.created_at, r.updated_at
|
||||
FROM repositories r
|
||||
WHERE r.owner_user_id = ?
|
||||
OR r.visibility = 'public'
|
||||
OR EXISTS (SELECT 1 FROM repository_members m WHERE m.repository_id = r.id AND m.user_id = ?)
|
||||
ORDER BY r.name`, userID, userID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list repos: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []*models.Repository
|
||||
for rows.Next() {
|
||||
r := &models.Repository{}
|
||||
if err := rows.Scan(&r.ID, &r.Name, &r.OwnerUserID, &r.Visibility, &r.Description, &r.CreatedAt, &r.UpdatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, r)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// UpdateRepository mutates a repository's name, visibility, and description.
|
||||
// Empty strings leave the field unchanged.
|
||||
func (s *Store) UpdateRepository(ctx context.Context, id, name string, visibility *models.Visibility, description *string) error {
|
||||
now := s.now()
|
||||
if name != "" {
|
||||
if _, err := s.exec(ctx, `UPDATE repositories SET name = ?, updated_at = ? WHERE id = ?`, name, now, id); err != nil {
|
||||
return fmt.Errorf("update repo name: %w", err)
|
||||
}
|
||||
}
|
||||
if visibility != nil {
|
||||
if _, err := s.exec(ctx, `UPDATE repositories SET visibility = ?, updated_at = ? WHERE id = ?`, string(*visibility), now, id); err != nil {
|
||||
return fmt.Errorf("update repo visibility: %w", err)
|
||||
}
|
||||
}
|
||||
if description != nil {
|
||||
if _, err := s.exec(ctx, `UPDATE repositories SET description = ?, updated_at = ? WHERE id = ?`, *description, now, id); err != nil {
|
||||
return fmt.Errorf("update repo description: %w", err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteRepository deletes a repository and all its child rows (cascade).
|
||||
// The caller must have already handled blob ref-count cleanup.
|
||||
func (s *Store) DeleteRepository(ctx context.Context, id string) error {
|
||||
if _, err := s.exec(ctx, `DELETE FROM repositories WHERE id = ?`, id); err != nil {
|
||||
return fmt.Errorf("delete repository: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetRepositoryByName returns a repository by its unique name.
|
||||
func (s *Store) GetRepositoryByName(ctx context.Context, name string) (*models.Repository, error) {
|
||||
row := s.db.QueryRowContext(ctx,
|
||||
`SELECT id, name, owner_user_id, visibility, description, created_at, updated_at
|
||||
FROM repositories WHERE name = ?`, name)
|
||||
return scanRepo(row)
|
||||
}
|
||||
|
||||
// GetRepositoryByID returns a repository by id.
|
||||
func (s *Store) GetRepositoryByID(ctx context.Context, id string) (*models.Repository, error) {
|
||||
row := s.db.QueryRowContext(ctx,
|
||||
`SELECT id, name, owner_user_id, visibility, description, created_at, updated_at
|
||||
FROM repositories WHERE id = ?`, id)
|
||||
return scanRepo(row)
|
||||
}
|
||||
|
||||
func scanRepo(row *sql.Row) (*models.Repository, error) {
|
||||
r := &models.Repository{}
|
||||
err := row.Scan(&r.ID, &r.Name, &r.OwnerUserID, &r.Visibility, &r.Description, &r.CreatedAt, &r.UpdatedAt)
|
||||
if isErrNoRows(err) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return r, nil
|
||||
}
|
||||
|
||||
// GetMemberAccess returns the access level granted to userID on repoID, or
|
||||
// ("", false) if the user is not an explicit member.
|
||||
func (s *Store) GetMemberAccess(ctx context.Context, repoID, userID string) (models.Access, bool, error) {
|
||||
row := s.db.QueryRowContext(ctx,
|
||||
`SELECT access FROM repository_members WHERE repository_id = ? AND user_id = ?`,
|
||||
repoID, userID)
|
||||
var access string
|
||||
err := row.Scan(&access)
|
||||
if isErrNoRows(err) {
|
||||
return "", false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return "", false, fmt.Errorf("get member access: %w", err)
|
||||
}
|
||||
return models.Access(access), true, nil
|
||||
}
|
||||
@@ -0,0 +1,253 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
func TestCreateRepository(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
|
||||
r := s.createRepo(t, ctx, "myrepo", u.ID, models.VisibilityPublic)
|
||||
if r.Name != "myrepo" || r.OwnerUserID != u.ID || r.Visibility != models.VisibilityPublic {
|
||||
t.Fatalf("repo = %+v", r)
|
||||
}
|
||||
if r.ID == "" || r.CreatedAt == "" {
|
||||
t.Fatal("id/created_at should be set")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetRepositoryByNameAndID(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
r := s.createRepo(t, ctx, "myrepo", u.ID, models.VisibilityPublic)
|
||||
|
||||
byName, err := s.GetRepositoryByName(ctx, "myrepo")
|
||||
if err != nil {
|
||||
t.Fatalf("get by name: %v", err)
|
||||
}
|
||||
if byName.ID != r.ID {
|
||||
t.Fatalf("id mismatch")
|
||||
}
|
||||
byID, err := s.GetRepositoryByID(ctx, r.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("get by id: %v", err)
|
||||
}
|
||||
if byID.Name != "myrepo" {
|
||||
t.Fatalf("name = %q", byID.Name)
|
||||
}
|
||||
if _, err := s.GetRepositoryByName(ctx, "missing"); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListReposVisible(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
alice := s.createUser(t, ctx, "alice", "p")
|
||||
bob := s.createUser(t, ctx, "bob", "p")
|
||||
carol := s.createUser(t, ctx, "carol", "p")
|
||||
|
||||
// alice owns: pub-own, priv-own
|
||||
pubOwn := s.createRepo(t, ctx, "pub-own", alice.ID, models.VisibilityPublic)
|
||||
privOwn := s.createRepo(t, ctx, "priv-own", alice.ID, models.VisibilityPrivate)
|
||||
// bob owns: pub-bob, priv-bob, and grants carol read on priv-bob
|
||||
pubBob := s.createRepo(t, ctx, "pub-bob", bob.ID, models.VisibilityPublic)
|
||||
privBob := s.createRepo(t, ctx, "priv-bob", bob.ID, models.VisibilityPrivate)
|
||||
_ = pubOwn
|
||||
_ = privOwn
|
||||
_ = pubBob
|
||||
if err := s.AddMember(ctx, privBob.ID, carol.ID, models.AccessRead); err != nil {
|
||||
t.Fatalf("add member: %v", err)
|
||||
}
|
||||
|
||||
// Alice sees: her two + bob's public. NOT bob's private.
|
||||
aliceRepos, err := s.ListReposVisible(ctx, alice.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("alice list: %v", err)
|
||||
}
|
||||
if len(aliceRepos) != 3 {
|
||||
t.Fatalf("alice should see 3 repos, got %d", len(aliceRepos))
|
||||
}
|
||||
|
||||
// Carol sees: pub-own, pub-bob (both public) + priv-bob (member). NOT priv-own.
|
||||
carolRepos, err := s.ListReposVisible(ctx, carol.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("carol list: %v", err)
|
||||
}
|
||||
if len(carolRepos) != 3 {
|
||||
t.Fatalf("carol should see 3 repos, got %d", len(carolRepos))
|
||||
}
|
||||
// Verify priv-bob is among carol's repos.
|
||||
foundPrivBob := false
|
||||
for _, r := range carolRepos {
|
||||
if r.ID == privBob.ID {
|
||||
foundPrivBob = true
|
||||
}
|
||||
if r.ID == privOwn.ID {
|
||||
t.Fatal("carol should not see alice's private repo")
|
||||
}
|
||||
}
|
||||
if !foundPrivBob {
|
||||
t.Fatal("carol should see priv-bob as a member")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateRepository(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
r := s.createRepo(t, ctx, "myrepo", u.ID, models.VisibilityPublic)
|
||||
|
||||
// Update name only.
|
||||
if err := s.UpdateRepository(ctx, r.ID, "newname", nil, nil); err != nil {
|
||||
t.Fatalf("update name: %v", err)
|
||||
}
|
||||
got, _ := s.GetRepositoryByID(ctx, r.ID)
|
||||
if got.Name != "newname" || got.Visibility != models.VisibilityPublic {
|
||||
t.Fatalf("name=%q vis=%s", got.Name, got.Visibility)
|
||||
}
|
||||
|
||||
// Update visibility only.
|
||||
priv := models.VisibilityPrivate
|
||||
_ = s.UpdateRepository(ctx, r.ID, "", &priv, nil)
|
||||
got, _ = s.GetRepositoryByID(ctx, r.ID)
|
||||
if got.Visibility != models.VisibilityPrivate || got.Name != "newname" {
|
||||
t.Fatalf("vis=%s name=%q", got.Visibility, got.Name)
|
||||
}
|
||||
|
||||
// Update description only.
|
||||
desc := "a repo"
|
||||
_ = s.UpdateRepository(ctx, r.ID, "", nil, &desc)
|
||||
got, _ = s.GetRepositoryByID(ctx, r.ID)
|
||||
if got.Description != "a repo" {
|
||||
t.Fatalf("desc=%q", got.Description)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteRepository(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
r := s.createRepo(t, ctx, "myrepo", u.ID, models.VisibilityPublic)
|
||||
|
||||
if err := s.DeleteRepository(ctx, r.ID); err != nil {
|
||||
t.Fatalf("delete: %v", err)
|
||||
}
|
||||
if _, err := s.GetRepositoryByID(ctx, r.ID); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound after delete, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// --- members ---
|
||||
|
||||
func TestAddMember_Upsert(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
alice := s.createUser(t, ctx, "alice", "p")
|
||||
bob := s.createUser(t, ctx, "bob", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", alice.ID, models.VisibilityPrivate)
|
||||
|
||||
// Initial grant: read.
|
||||
if err := s.AddMember(ctx, repo.ID, bob.ID, models.AccessRead); err != nil {
|
||||
t.Fatalf("add: %v", err)
|
||||
}
|
||||
access, ok, err := s.GetMemberAccess(ctx, repo.ID, bob.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("get: %v", err)
|
||||
}
|
||||
if !ok || access != models.AccessRead {
|
||||
t.Fatalf("expected read grant, got ok=%v access=%s", ok, access)
|
||||
}
|
||||
|
||||
// Upsert to write.
|
||||
if err := s.AddMember(ctx, repo.ID, bob.ID, models.AccessWrite); err != nil {
|
||||
t.Fatalf("upsert: %v", err)
|
||||
}
|
||||
access, ok, _ = s.GetMemberAccess(ctx, repo.ID, bob.ID)
|
||||
if !ok || access != models.AccessWrite {
|
||||
t.Fatalf("expected write after upsert, got %s", access)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetMemberAccess_NonMember(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
alice := s.createUser(t, ctx, "alice", "p")
|
||||
bob := s.createUser(t, ctx, "bob", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", alice.ID, models.VisibilityPrivate)
|
||||
|
||||
_, ok, err := s.GetMemberAccess(ctx, repo.ID, bob.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("get: %v", err)
|
||||
}
|
||||
if ok {
|
||||
t.Fatal("non-member should return ok=false")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateMemberAccess_NotFound(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
alice := s.createUser(t, ctx, "alice", "p")
|
||||
bob := s.createUser(t, ctx, "bob", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", alice.ID, models.VisibilityPrivate)
|
||||
|
||||
if err := s.UpdateMemberAccess(ctx, repo.ID, bob.ID, models.AccessRead); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("update non-member should be ErrNotFound, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveMember(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
alice := s.createUser(t, ctx, "alice", "p")
|
||||
bob := s.createUser(t, ctx, "bob", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", alice.ID, models.VisibilityPrivate)
|
||||
|
||||
_ = s.AddMember(ctx, repo.ID, bob.ID, models.AccessRead)
|
||||
if err := s.RemoveMember(ctx, repo.ID, bob.ID); err != nil {
|
||||
t.Fatalf("remove: %v", err)
|
||||
}
|
||||
_, ok, _ := s.GetMemberAccess(ctx, repo.ID, bob.ID)
|
||||
if ok {
|
||||
t.Fatal("member should be gone after remove")
|
||||
}
|
||||
// Removing again returns ErrNotFound.
|
||||
if err := s.RemoveMember(ctx, repo.ID, bob.ID); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("second remove should be ErrNotFound, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListMembers(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
alice := s.createUser(t, ctx, "alice", "p")
|
||||
bob := s.createUser(t, ctx, "bob", "p")
|
||||
carol := s.createUser(t, ctx, "carol", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", alice.ID, models.VisibilityPrivate)
|
||||
|
||||
_ = s.AddMember(ctx, repo.ID, carol.ID, models.AccessRead)
|
||||
_ = s.AddMember(ctx, repo.ID, bob.ID, models.AccessWrite)
|
||||
|
||||
members, err := s.ListMembers(ctx, repo.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("list: %v", err)
|
||||
}
|
||||
if len(members) != 2 {
|
||||
t.Fatalf("expected 2 members, got %d", len(members))
|
||||
}
|
||||
// Ordered by username: bob, carol.
|
||||
if members[0].User.Username != "bob" || members[0].Access != models.AccessWrite {
|
||||
t.Fatalf("member[0] = %+v", members[0])
|
||||
}
|
||||
if members[1].User.Username != "carol" || members[1].Access != models.AccessRead {
|
||||
t.Fatalf("member[1] = %+v", members[1])
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
// Package store provides the data-access layer for urapt-server: typed
|
||||
// methods over the SQLite database backing all domain objects (users, tokens,
|
||||
// repositories, distributions, components, architectures, packages, blobs,
|
||||
// gpg keys, audit log).
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// Store is the entrypoint to the data-access layer. All methods are safe for
|
||||
// concurrent use; the underlying *sql.DB is configured with a single writer
|
||||
// connection (see shared/db).
|
||||
type Store struct {
|
||||
db *sql.DB
|
||||
now func() string
|
||||
}
|
||||
|
||||
// New constructs a Store wrapping db.
|
||||
func New(db *sql.DB) *Store {
|
||||
return &Store{db: db, now: nowISO}
|
||||
}
|
||||
|
||||
// DB returns the underlying database (used by app for raw queries if needed).
|
||||
func (s *Store) DB() *sql.DB { return s.db }
|
||||
|
||||
// Now returns the current timestamp in the urapt canonical form.
|
||||
func (s *Store) Now() string { return s.now() }
|
||||
|
||||
// newID returns a fresh UUIDv4 string.
|
||||
func newID() string { return uuid.NewString() }
|
||||
|
||||
// nowISO returns the current UTC time in RFC3339 form.
|
||||
func nowISO() string { return time.Now().UTC().Format(time.RFC3339Nano) }
|
||||
|
||||
// exec is a small helper for ExecContext with a context.
|
||||
func (s *Store) exec(ctx context.Context, query string, args ...any) (sql.Result, error) {
|
||||
return s.db.ExecContext(ctx, query, args...)
|
||||
}
|
||||
|
||||
// ErrNotFound is returned by Get-style methods when no row matches.
|
||||
var ErrNotFound = fmt.Errorf("not found")
|
||||
|
||||
// isErrNoRows returns true if err is sql.ErrNoRows.
|
||||
func isErrNoRows(err error) bool { return err == sql.ErrNoRows }
|
||||
@@ -0,0 +1,87 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"urapt/shared/crypto"
|
||||
"urapt/shared/db"
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// newTestStore opens a fresh migrated SQLite database in a per-test temp
|
||||
// directory and returns a Store over it. The database is closed automatically
|
||||
// when the test finishes.
|
||||
func newTestStore(t *testing.T) *Store {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
database, err := db.Open(filepath.Join(dir, "test.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("db open: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { database.Close() })
|
||||
return New(database)
|
||||
}
|
||||
|
||||
// createUser is a test helper that inserts a user with a bcrypt-hashed
|
||||
// password and returns it. The password is hashed for realism so that
|
||||
// password-verification paths can be exercised if needed.
|
||||
func (s *Store) createUser(t *testing.T, ctx context.Context, username, password string) *models.User {
|
||||
t.Helper()
|
||||
hash, err := crypto.HashPassword(password)
|
||||
if err != nil {
|
||||
t.Fatalf("hash password: %v", err)
|
||||
}
|
||||
u, admin, err := s.CreateUser(ctx, username, hash)
|
||||
if err != nil {
|
||||
t.Fatalf("create user %q: %v", username, err)
|
||||
}
|
||||
t.Logf("created user %q (admin=%v)", u.Username, admin)
|
||||
return u
|
||||
}
|
||||
|
||||
// createToken is a test helper that mints a real API token (with hash and
|
||||
// prefix) for userID and returns the plaintext token plus the stored row.
|
||||
func (s *Store) createToken(t *testing.T, ctx context.Context, userID, name string) (string, *models.APIToken) {
|
||||
t.Helper()
|
||||
tok, hash, prefix, err := crypto.GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
stored, err := s.CreateToken(ctx, userID, name, prefix, hash)
|
||||
if err != nil {
|
||||
t.Fatalf("create token: %v", err)
|
||||
}
|
||||
return tok, stored
|
||||
}
|
||||
|
||||
// createRepo is a test helper that creates a repository owned by userID.
|
||||
func (s *Store) createRepo(t *testing.T, ctx context.Context, name, ownerID string, vis models.Visibility) *models.Repository {
|
||||
t.Helper()
|
||||
r, err := s.CreateRepository(ctx, name, ownerID, vis, "")
|
||||
if err != nil {
|
||||
t.Fatalf("create repo %q: %v", name, err)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// createDistro is a test helper that creates a distribution within repoID.
|
||||
func (s *Store) createDistro(t *testing.T, ctx context.Context, repoID, name string) *models.Distribution {
|
||||
t.Helper()
|
||||
d, err := s.CreateDistribution(ctx, repoID, name)
|
||||
if err != nil {
|
||||
t.Fatalf("create distro %q: %v", name, err)
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
// createComponent is a test helper that creates a component within distroID.
|
||||
func (s *Store) createComponent(t *testing.T, ctx context.Context, distroID, name string) *models.Component {
|
||||
t.Helper()
|
||||
c, err := s.CreateComponent(ctx, distroID, name)
|
||||
if err != nil {
|
||||
t.Fatalf("create component %q: %v", name, err)
|
||||
}
|
||||
return c
|
||||
}
|
||||
@@ -0,0 +1,190 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// --- distributions ---
|
||||
|
||||
// CreateDistribution adds a distribution (suite) to a repository.
|
||||
func (s *Store) CreateDistribution(ctx context.Context, repoID, name string) (*models.Distribution, error) {
|
||||
id := newID()
|
||||
now := s.now()
|
||||
_, err := s.exec(ctx, `INSERT INTO distributions (id, repository_id, name, created_at) VALUES (?, ?, ?, ?)`,
|
||||
id, repoID, name, now)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("insert distribution: %w", err)
|
||||
}
|
||||
return &models.Distribution{ID: id, RepositoryID: repoID, Name: name, CreatedAt: now}, nil
|
||||
}
|
||||
|
||||
// GetDistributionByName returns a distribution by name within a repository.
|
||||
func (s *Store) GetDistributionByName(ctx context.Context, repoID, name string) (*models.Distribution, error) {
|
||||
row := s.db.QueryRowContext(ctx,
|
||||
`SELECT id, repository_id, name, created_at FROM distributions WHERE repository_id = ? AND name = ?`, repoID, name)
|
||||
d := &models.Distribution{}
|
||||
err := row.Scan(&d.ID, &d.RepositoryID, &d.Name, &d.CreatedAt)
|
||||
if isErrNoRows(err) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return d, err
|
||||
}
|
||||
|
||||
// ListDistributions returns all distributions in a repository.
|
||||
func (s *Store) ListDistributions(ctx context.Context, repoID string) ([]*models.Distribution, error) {
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, repository_id, name, created_at FROM distributions WHERE repository_id = ? ORDER BY name`, repoID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list distributions: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []*models.Distribution
|
||||
for rows.Next() {
|
||||
d := &models.Distribution{}
|
||||
if err := rows.Scan(&d.ID, &d.RepositoryID, &d.Name, &d.CreatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, d)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// DeleteDistribution removes a distribution and cascades to components,
|
||||
// architectures, and packages.
|
||||
func (s *Store) DeleteDistribution(ctx context.Context, repoID, name string) error {
|
||||
res, err := s.exec(ctx, `DELETE FROM distributions WHERE repository_id = ? AND name = ?`, repoID, name)
|
||||
if err != nil {
|
||||
return fmt.Errorf("delete distribution: %w", err)
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return ErrNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// --- components ---
|
||||
|
||||
// CreateComponent adds a component to a distribution.
|
||||
func (s *Store) CreateComponent(ctx context.Context, distroID, name string) (*models.Component, error) {
|
||||
id := newID()
|
||||
now := s.now()
|
||||
_, err := s.exec(ctx, `INSERT INTO components (id, distribution_id, name, created_at) VALUES (?, ?, ?, ?)`,
|
||||
id, distroID, name, now)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("insert component: %w", err)
|
||||
}
|
||||
return &models.Component{ID: id, DistributionID: distroID, Name: name, CreatedAt: now}, nil
|
||||
}
|
||||
|
||||
// GetComponentByName returns a component by name within a distribution.
|
||||
func (s *Store) GetComponentByName(ctx context.Context, distroID, name string) (*models.Component, error) {
|
||||
row := s.db.QueryRowContext(ctx,
|
||||
`SELECT id, distribution_id, name, created_at FROM components WHERE distribution_id = ? AND name = ?`, distroID, name)
|
||||
c := &models.Component{}
|
||||
err := row.Scan(&c.ID, &c.DistributionID, &c.Name, &c.CreatedAt)
|
||||
if isErrNoRows(err) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return c, err
|
||||
}
|
||||
|
||||
// ListComponents returns all components in a distribution.
|
||||
func (s *Store) ListComponents(ctx context.Context, distroID string) ([]*models.Component, error) {
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, distribution_id, name, created_at FROM components WHERE distribution_id = ? ORDER BY name`, distroID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list components: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []*models.Component
|
||||
for rows.Next() {
|
||||
c := &models.Component{}
|
||||
if err := rows.Scan(&c.ID, &c.DistributionID, &c.Name, &c.CreatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, c)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// DeleteComponent removes a component. The schema blocks deletion while
|
||||
// packages reference it (ON DELETE RESTRICT); the handler checks first.
|
||||
func (s *Store) DeleteComponent(ctx context.Context, distroID, name string) error {
|
||||
res, err := s.exec(ctx, `DELETE FROM components WHERE distribution_id = ? AND name = ?`, distroID, name)
|
||||
if err != nil {
|
||||
return fmt.Errorf("delete component: %w", err)
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return ErrNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// CountComponentsByDistro reports how many components a distribution has.
|
||||
func (s *Store) CountComponentsByDistro(ctx context.Context, distroID string) (int, error) {
|
||||
var n int
|
||||
err := s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM components WHERE distribution_id = ?`, distroID).Scan(&n)
|
||||
return n, err
|
||||
}
|
||||
|
||||
// --- architectures ---
|
||||
|
||||
// CreateArchitecture adds an architecture to a distribution.
|
||||
func (s *Store) CreateArchitecture(ctx context.Context, distroID, name string) (*models.Architecture, error) {
|
||||
id := newID()
|
||||
now := s.now()
|
||||
_, err := s.exec(ctx, `INSERT INTO architectures (id, distribution_id, name, created_at) VALUES (?, ?, ?, ?)`,
|
||||
id, distroID, name, now)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("insert architecture: %w", err)
|
||||
}
|
||||
return &models.Architecture{ID: id, DistributionID: distroID, Name: name, CreatedAt: now}, nil
|
||||
}
|
||||
|
||||
// ListArchitectures returns all architectures in a distribution.
|
||||
func (s *Store) ListArchitectures(ctx context.Context, distroID string) ([]*models.Architecture, error) {
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, distribution_id, name, created_at FROM architectures WHERE distribution_id = ? ORDER BY name`, distroID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list architectures: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []*models.Architecture
|
||||
for rows.Next() {
|
||||
a := &models.Architecture{}
|
||||
if err := rows.Scan(&a.ID, &a.DistributionID, &a.Name, &a.CreatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, a)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// DeleteArchitecture removes an architecture from a distribution.
|
||||
func (s *Store) DeleteArchitecture(ctx context.Context, distroID, name string) error {
|
||||
res, err := s.exec(ctx, `DELETE FROM architectures WHERE distribution_id = ? AND name = ?`, distroID, name)
|
||||
if err != nil {
|
||||
return fmt.Errorf("delete architecture: %w", err)
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return ErrNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// HasArchitecture reports whether a distribution declares the given arch.
|
||||
func (s *Store) HasArchitecture(ctx context.Context, distroID, name string) (bool, error) {
|
||||
var n int
|
||||
err := s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM architectures WHERE distribution_id = ? AND name = ?`, distroID, name).Scan(&n)
|
||||
if err == sql.ErrNoRows {
|
||||
return false, nil
|
||||
}
|
||||
return n > 0, err
|
||||
}
|
||||
@@ -0,0 +1,204 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// --- distributions ---
|
||||
|
||||
func TestCreateDistributionAndGetByName(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, "public")
|
||||
|
||||
d := s.createDistro(t, ctx, repo.ID, "stable")
|
||||
if d.Name != "stable" || d.RepositoryID != repo.ID {
|
||||
t.Fatalf("distro = %+v", d)
|
||||
}
|
||||
|
||||
got, err := s.GetDistributionByName(ctx, repo.ID, "stable")
|
||||
if err != nil {
|
||||
t.Fatalf("get: %v", err)
|
||||
}
|
||||
if got.ID != d.ID {
|
||||
t.Fatal("id mismatch")
|
||||
}
|
||||
if _, err := s.GetDistributionByName(ctx, repo.ID, "missing"); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListDistributions(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, "public")
|
||||
s.createDistro(t, ctx, repo.ID, "stable")
|
||||
s.createDistro(t, ctx, repo.ID, "unstable")
|
||||
s.createDistro(t, ctx, repo.ID, "oldstable")
|
||||
|
||||
ds, err := s.ListDistributions(ctx, repo.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("list: %v", err)
|
||||
}
|
||||
if len(ds) != 3 {
|
||||
t.Fatalf("expected 3 distros, got %d", len(ds))
|
||||
}
|
||||
// Ordered by name.
|
||||
if ds[0].Name != "oldstable" || ds[1].Name != "stable" || ds[2].Name != "unstable" {
|
||||
names := []string{ds[0].Name, ds[1].Name, ds[2].Name}
|
||||
t.Fatalf("expected sorted, got %v", names)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteDistribution(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, "public")
|
||||
s.createDistro(t, ctx, repo.ID, "stable")
|
||||
|
||||
if err := s.DeleteDistribution(ctx, repo.ID, "stable"); err != nil {
|
||||
t.Fatalf("delete: %v", err)
|
||||
}
|
||||
if _, err := s.GetDistributionByName(ctx, repo.ID, "stable"); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound after delete, got %v", err)
|
||||
}
|
||||
if err := s.DeleteDistribution(ctx, repo.ID, "stable"); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("second delete should be ErrNotFound, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// --- components ---
|
||||
|
||||
func TestCreateComponentAndList(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, "public")
|
||||
d := s.createDistro(t, ctx, repo.ID, "stable")
|
||||
|
||||
s.createComponent(t, ctx, d.ID, "main")
|
||||
s.createComponent(t, ctx, d.ID, "contrib")
|
||||
s.createComponent(t, ctx, d.ID, "non-free")
|
||||
|
||||
cs, err := s.ListComponents(ctx, d.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("list: %v", err)
|
||||
}
|
||||
if len(cs) != 3 {
|
||||
t.Fatalf("expected 3 components, got %d", len(cs))
|
||||
}
|
||||
if cs[0].Name != "contrib" {
|
||||
t.Fatalf("expected sorted; first = %q", cs[0].Name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteComponent(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, "public")
|
||||
d := s.createDistro(t, ctx, repo.ID, "stable")
|
||||
s.createComponent(t, ctx, d.ID, "main")
|
||||
|
||||
if err := s.DeleteComponent(ctx, d.ID, "main"); err != nil {
|
||||
t.Fatalf("delete: %v", err)
|
||||
}
|
||||
if err := s.DeleteComponent(ctx, d.ID, "main"); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("second delete should be ErrNotFound, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCountComponentsByDistro(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, "public")
|
||||
d := s.createDistro(t, ctx, repo.ID, "stable")
|
||||
|
||||
n, err := s.CountComponentsByDistro(ctx, d.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("count: %v", err)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Fatalf("expected 0, got %d", n)
|
||||
}
|
||||
s.createComponent(t, ctx, d.ID, "main")
|
||||
s.createComponent(t, ctx, d.ID, "contrib")
|
||||
n, _ = s.CountComponentsByDistro(ctx, d.ID)
|
||||
if n != 2 {
|
||||
t.Fatalf("expected 2, got %d", n)
|
||||
}
|
||||
}
|
||||
|
||||
// --- architectures ---
|
||||
|
||||
func TestCreateArchitectureAndList(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, "public")
|
||||
d := s.createDistro(t, ctx, repo.ID, "stable")
|
||||
|
||||
if _, err := s.CreateArchitecture(ctx, d.ID, "amd64"); err != nil {
|
||||
t.Fatalf("create amd64: %v", err)
|
||||
}
|
||||
if _, err := s.CreateArchitecture(ctx, d.ID, "arm64"); err != nil {
|
||||
t.Fatalf("create arm64: %v", err)
|
||||
}
|
||||
|
||||
arches, err := s.ListArchitectures(ctx, d.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("list: %v", err)
|
||||
}
|
||||
if len(arches) != 2 {
|
||||
t.Fatalf("expected 2 arches, got %d", len(arches))
|
||||
}
|
||||
if arches[0].Name != "amd64" || arches[1].Name != "arm64" {
|
||||
t.Fatalf("expected sorted, got %q %q", arches[0].Name, arches[1].Name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteArchitecture(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, "public")
|
||||
d := s.createDistro(t, ctx, repo.ID, "stable")
|
||||
_, _ = s.CreateArchitecture(ctx, d.ID, "amd64")
|
||||
|
||||
if err := s.DeleteArchitecture(ctx, d.ID, "amd64"); err != nil {
|
||||
t.Fatalf("delete: %v", err)
|
||||
}
|
||||
if err := s.DeleteArchitecture(ctx, d.ID, "amd64"); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("second delete should be ErrNotFound, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHasArchitecture(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
repo := s.createRepo(t, ctx, "repo", u.ID, "public")
|
||||
d := s.createDistro(t, ctx, repo.ID, "stable")
|
||||
_, _ = s.CreateArchitecture(ctx, d.ID, "amd64")
|
||||
|
||||
has, err := s.HasArchitecture(ctx, d.ID, "amd64")
|
||||
if err != nil {
|
||||
t.Fatalf("has amd64: %v", err)
|
||||
}
|
||||
if !has {
|
||||
t.Fatal("expected has=true for amd64")
|
||||
}
|
||||
has, err = s.HasArchitecture(ctx, d.ID, "arm64")
|
||||
if err != nil {
|
||||
t.Fatalf("has arm64: %v", err)
|
||||
}
|
||||
if has {
|
||||
t.Fatal("expected has=false for arm64")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// CreateToken inserts a new API token row. The plaintext token is NOT stored;
|
||||
// only its SHA-256 hash and display prefix are.
|
||||
func (s *Store) CreateToken(ctx context.Context, userID, name, prefix, tokenHash string) (*models.APIToken, error) {
|
||||
id := newID()
|
||||
now := s.now()
|
||||
_, err := s.exec(ctx, `INSERT INTO api_tokens (id, user_id, name, prefix, token_hash, created_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?)`, id, userID, name, prefix, tokenHash, now)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("insert token: %w", err)
|
||||
}
|
||||
return &models.APIToken{
|
||||
ID: id, UserID: userID, Name: name, Prefix: prefix, CreatedAt: now,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetTokenByHash returns the active (non-revoked) token with the given hash.
|
||||
func (s *Store) GetTokenByHash(ctx context.Context, hash string) (*models.APIToken, error) {
|
||||
row := s.db.QueryRowContext(ctx,
|
||||
`SELECT id, user_id, name, prefix, created_at, last_used_at, revoked_at
|
||||
FROM api_tokens WHERE token_hash = ? AND revoked_at IS NULL`, hash)
|
||||
t := &models.APIToken{}
|
||||
var lastUsed, revoked sql.NullString
|
||||
err := row.Scan(&t.ID, &t.UserID, &t.Name, &t.Prefix, &t.CreatedAt, &lastUsed, &revoked)
|
||||
if isErrNoRows(err) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if lastUsed.Valid {
|
||||
v := lastUsed.String
|
||||
t.LastUsedAt = &v
|
||||
}
|
||||
if revoked.Valid {
|
||||
v := revoked.String
|
||||
t.RevokedAt = &v
|
||||
}
|
||||
return t, nil
|
||||
}
|
||||
|
||||
// TouchToken updates last_used_at for a token.
|
||||
func (s *Store) TouchToken(ctx context.Context, id string) error {
|
||||
_, err := s.exec(ctx, `UPDATE api_tokens SET last_used_at = ? WHERE id = ?`, s.now(), id)
|
||||
return err
|
||||
}
|
||||
|
||||
// ListTokens returns all tokens for a user (including revoked).
|
||||
func (s *Store) ListTokens(ctx context.Context, userID string) ([]*models.APIToken, error) {
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, user_id, name, prefix, created_at, last_used_at, revoked_at
|
||||
FROM api_tokens WHERE user_id = ? ORDER BY created_at`, userID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list tokens: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []*models.APIToken
|
||||
for rows.Next() {
|
||||
t := &models.APIToken{}
|
||||
var lastUsed, revoked sql.NullString
|
||||
if err := rows.Scan(&t.ID, &t.UserID, &t.Name, &t.Prefix, &t.CreatedAt, &lastUsed, &revoked); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if lastUsed.Valid {
|
||||
v := lastUsed.String
|
||||
t.LastUsedAt = &v
|
||||
}
|
||||
if revoked.Valid {
|
||||
v := revoked.String
|
||||
t.RevokedAt = &v
|
||||
}
|
||||
out = append(out, t)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// RevokeToken marks the given token revoked. It must belong to userID.
|
||||
func (s *Store) RevokeToken(ctx context.Context, userID, tokenID string) error {
|
||||
res, err := s.exec(ctx, `UPDATE api_tokens SET revoked_at = ? WHERE id = ? AND user_id = ?`,
|
||||
s.now(), tokenID, userID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("revoke token: %w", err)
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return ErrNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RevokeTokenByHash marks the token with the given hash revoked (used for logout).
|
||||
func (s *Store) RevokeTokenByHash(ctx context.Context, hash string) error {
|
||||
_, err := s.exec(ctx, `UPDATE api_tokens SET revoked_at = ? WHERE token_hash = ? AND revoked_at IS NULL`,
|
||||
s.now(), hash)
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,127 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"urapt/shared/models"
|
||||
)
|
||||
|
||||
// CreateUser inserts a new user. If the users table is empty, the user is made
|
||||
// an admin (the bootstrap-admin rule).
|
||||
func (s *Store) CreateUser(ctx context.Context, username, passwordHash string) (*models.User, bool, error) {
|
||||
username = strings.TrimSpace(username)
|
||||
lc := strings.ToLower(username)
|
||||
id := newID()
|
||||
now := s.now()
|
||||
|
||||
var admin bool
|
||||
var count int
|
||||
if err := s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM users`).Scan(&count); err != nil {
|
||||
return nil, false, fmt.Errorf("count users: %w", err)
|
||||
}
|
||||
admin = count == 0
|
||||
|
||||
_, err := s.exec(ctx, `INSERT INTO users (id, username, username_lc, password_hash, is_admin, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)`,
|
||||
id, username, lc, passwordHash, boolToInt(admin), now, now)
|
||||
if err != nil {
|
||||
return nil, false, fmt.Errorf("insert user: %w", err)
|
||||
}
|
||||
return &models.User{
|
||||
ID: id, Username: username, IsAdmin: admin, CreatedAt: now, UpdatedAt: now,
|
||||
}, admin, nil
|
||||
}
|
||||
|
||||
// GetUserByID returns the user with the given id.
|
||||
func (s *Store) GetUserByID(ctx context.Context, id string) (*models.User, error) {
|
||||
row := s.db.QueryRowContext(ctx,
|
||||
`SELECT id, username, is_admin, created_at, updated_at FROM users WHERE id = ?`, id)
|
||||
return scanUser(row)
|
||||
}
|
||||
|
||||
// GetUserByUsername returns the user with the given (case-insensitive) username.
|
||||
func (s *Store) GetUserByUsername(ctx context.Context, username string) (*models.User, error) {
|
||||
row := s.db.QueryRowContext(ctx,
|
||||
`SELECT id, username, is_admin, created_at, updated_at FROM users WHERE username_lc = ?`,
|
||||
strings.ToLower(strings.TrimSpace(username)))
|
||||
return scanUser(row)
|
||||
}
|
||||
|
||||
// ListUsers returns all users ordered by username.
|
||||
func (s *Store) ListUsers(ctx context.Context) ([]*models.User, error) {
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, username, is_admin, created_at, updated_at FROM users ORDER BY username_lc`)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list users: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []*models.User
|
||||
for rows.Next() {
|
||||
u, err := scanUserRows(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, u)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// GetUserPasswordHash returns the stored bcrypt hash for a user.
|
||||
func (s *Store) GetUserPasswordHash(ctx context.Context, id string) (string, error) {
|
||||
var hash string
|
||||
err := s.db.QueryRowContext(ctx, `SELECT password_hash FROM users WHERE id = ?`, id).Scan(&hash)
|
||||
if isErrNoRows(err) {
|
||||
return "", ErrNotFound
|
||||
}
|
||||
return hash, err
|
||||
}
|
||||
|
||||
// UpdateUser updates mutable fields. If isAdmin is nil, it is left unchanged.
|
||||
func (s *Store) UpdateUser(ctx context.Context, id string, isAdmin *bool) error {
|
||||
now := s.now()
|
||||
if isAdmin != nil {
|
||||
if _, err := s.exec(ctx, `UPDATE users SET is_admin = ?, updated_at = ? WHERE id = ?`,
|
||||
boolToInt(*isAdmin), now, id); err != nil {
|
||||
return fmt.Errorf("update user: %w", err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteUser removes a user. The caller should prevent self-deletion.
|
||||
func (s *Store) DeleteUser(ctx context.Context, id string) error {
|
||||
if _, err := s.exec(ctx, `DELETE FROM users WHERE id = ?`, id); err != nil {
|
||||
return fmt.Errorf("delete user: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func scanUser(row *sql.Row) (*models.User, error) {
|
||||
u := &models.User{}
|
||||
err := row.Scan(&u.ID, &u.Username, &u.IsAdmin, &u.CreatedAt, &u.UpdatedAt)
|
||||
if isErrNoRows(err) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return u, nil
|
||||
}
|
||||
|
||||
func scanUserRows(rows *sql.Rows) (*models.User, error) {
|
||||
u := &models.User{}
|
||||
if err := rows.Scan(&u.ID, &u.Username, &u.IsAdmin, &u.CreatedAt, &u.UpdatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return u, nil
|
||||
}
|
||||
|
||||
func boolToInt(b bool) int {
|
||||
if b {
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
@@ -0,0 +1,299 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"urapt/shared/crypto"
|
||||
)
|
||||
|
||||
func TestCreateUser_BootstrapAdmin(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
|
||||
// First user becomes admin.
|
||||
u1 := s.createUser(t, ctx, "alice", "pw1")
|
||||
if !u1.IsAdmin {
|
||||
t.Fatalf("first user should be admin, got is_admin=%v", u1.IsAdmin)
|
||||
}
|
||||
|
||||
// Subsequent users are not admin.
|
||||
u2 := s.createUser(t, ctx, "bob", "pw2")
|
||||
if u2.IsAdmin {
|
||||
t.Fatalf("second user should not be admin")
|
||||
}
|
||||
u3 := s.createUser(t, ctx, "carol", "pw3")
|
||||
if u3.IsAdmin {
|
||||
t.Fatalf("third user should not be admin")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateUser_DuplicateUsernameRejected(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
s.createUser(t, ctx, "alice", "pw1")
|
||||
if _, _, err := s.CreateUser(ctx, "alice", "hash"); err == nil {
|
||||
t.Fatal("expected error creating duplicate username")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetUserByID(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "pw1")
|
||||
|
||||
got, err := s.GetUserByID(ctx, u.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("get by id: %v", err)
|
||||
}
|
||||
if got.ID != u.ID || got.Username != "alice" {
|
||||
t.Fatalf("got %+v", got)
|
||||
}
|
||||
|
||||
if _, err := s.GetUserByID(ctx, "nope"); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetUserByUsername_CaseInsensitive(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
s.createUser(t, ctx, "Alice", "pw")
|
||||
|
||||
for _, q := range []string{"alice", "ALICE", "AlIcE"} {
|
||||
got, err := s.GetUserByUsername(ctx, q)
|
||||
if err != nil {
|
||||
t.Fatalf("lookup %q: %v", q, err)
|
||||
}
|
||||
if got.Username != "Alice" {
|
||||
t.Fatalf("expected original casing 'Alice', got %q", got.Username)
|
||||
}
|
||||
}
|
||||
if _, err := s.GetUserByUsername(ctx, "bob"); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound for missing user, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetUserPasswordHash(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "supersecret")
|
||||
|
||||
hash, err := s.GetUserPasswordHash(ctx, u.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("get hash: %v", err)
|
||||
}
|
||||
if !crypto.VerifyPassword(hash, "supersecret") {
|
||||
t.Fatal("bcrypt hash did not verify against original password")
|
||||
}
|
||||
if crypto.VerifyPassword(hash, "wrong") {
|
||||
t.Fatal("bcrypt hash verified against wrong password")
|
||||
}
|
||||
}
|
||||
|
||||
func TestListUsers(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
s.createUser(t, ctx, "carol", "p")
|
||||
s.createUser(t, ctx, "alice", "p")
|
||||
s.createUser(t, ctx, "bob", "p")
|
||||
|
||||
users, err := s.ListUsers(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("list: %v", err)
|
||||
}
|
||||
if len(users) != 3 {
|
||||
t.Fatalf("expected 3 users, got %d", len(users))
|
||||
}
|
||||
// Ordered by username_lc.
|
||||
if users[0].Username != "alice" || users[1].Username != "bob" || users[2].Username != "carol" {
|
||||
names := []string{users[0].Username, users[1].Username, users[2].Username}
|
||||
t.Fatalf("expected alphabetical order, got %v", names)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateUser(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
// First user is bootstrap admin; create a second non-admin user to test
|
||||
// promotion/demotion on.
|
||||
s.createUser(t, ctx, "admin", "p")
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
if u.IsAdmin {
|
||||
t.Fatal("non-bootstrap user should not be admin")
|
||||
}
|
||||
|
||||
admin := true
|
||||
if err := s.UpdateUser(ctx, u.ID, &admin); err != nil {
|
||||
t.Fatalf("update: %v", err)
|
||||
}
|
||||
got, _ := s.GetUserByID(ctx, u.ID)
|
||||
if !got.IsAdmin {
|
||||
t.Fatal("expected is_admin=true after update")
|
||||
}
|
||||
|
||||
off := false
|
||||
_ = s.UpdateUser(ctx, u.ID, &off)
|
||||
got, _ = s.GetUserByID(ctx, u.ID)
|
||||
if got.IsAdmin {
|
||||
t.Fatal("expected is_admin=false after demotion")
|
||||
}
|
||||
|
||||
// nil leaves it unchanged.
|
||||
_ = s.UpdateUser(ctx, u.ID, nil)
|
||||
got, _ = s.GetUserByID(ctx, u.ID)
|
||||
if got.IsAdmin {
|
||||
t.Fatal("nil isAdmin should leave it false")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteUser(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
|
||||
if err := s.DeleteUser(ctx, u.ID); err != nil {
|
||||
t.Fatalf("delete: %v", err)
|
||||
}
|
||||
if _, err := s.GetUserByID(ctx, u.ID); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("expected ErrNotFound after delete, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// --- tokens ---
|
||||
|
||||
func TestCreateTokenAndGetByHash(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
|
||||
plaintext, stored := s.createToken(t, ctx, u.ID, "laptop")
|
||||
if stored.Name != "laptop" {
|
||||
t.Fatalf("name = %q", stored.Name)
|
||||
}
|
||||
if stored.Prefix == "" {
|
||||
t.Fatal("prefix should be set")
|
||||
}
|
||||
|
||||
got, err := s.GetTokenByHash(ctx, crypto.HashToken(plaintext))
|
||||
if err != nil {
|
||||
t.Fatalf("get by hash: %v", err)
|
||||
}
|
||||
if got.ID != stored.ID {
|
||||
t.Fatalf("id mismatch: %s vs %s", got.ID, stored.ID)
|
||||
}
|
||||
if got.RevokedAt != nil {
|
||||
t.Fatal("fresh token should not be revoked")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetTokenByHash_RevokedExcluded(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
plaintext, stored := s.createToken(t, ctx, u.ID, "laptop")
|
||||
|
||||
if err := s.RevokeToken(ctx, u.ID, stored.ID); err != nil {
|
||||
t.Fatalf("revoke: %v", err)
|
||||
}
|
||||
if _, err := s.GetTokenByHash(ctx, crypto.HashToken(plaintext)); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("revoked token should not be returned, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRevokeToken_OwnershipEnforced(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
alice := s.createUser(t, ctx, "alice", "p")
|
||||
bob := s.createUser(t, ctx, "bob", "p")
|
||||
_, aliceToken := s.createToken(t, ctx, alice.ID, "alice-laptop")
|
||||
|
||||
// Bob cannot revoke Alice's token.
|
||||
if err := s.RevokeToken(ctx, bob.ID, aliceToken.ID); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("cross-user revoke should be ErrNotFound, got %v", err)
|
||||
}
|
||||
// Alice can revoke her own.
|
||||
if err := s.RevokeToken(ctx, alice.ID, aliceToken.ID); err != nil {
|
||||
t.Fatalf("own revoke: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRevokeTokenByHash(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
plaintext, _ := s.createToken(t, ctx, u.ID, "laptop")
|
||||
|
||||
if err := s.RevokeTokenByHash(ctx, crypto.HashToken(plaintext)); err != nil {
|
||||
t.Fatalf("revoke by hash: %v", err)
|
||||
}
|
||||
if _, err := s.GetTokenByHash(ctx, crypto.HashToken(plaintext)); !errors.Is(err, ErrNotFound) {
|
||||
t.Fatalf("revoked token should not be found, got %v", err)
|
||||
}
|
||||
// Revoking again is a no-op (no error, no rows affected).
|
||||
if err := s.RevokeTokenByHash(ctx, crypto.HashToken(plaintext)); err != nil {
|
||||
t.Fatalf("idempotent revoke: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListTokens_IncludesRevoked(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
_, t1 := s.createToken(t, ctx, u.ID, "a")
|
||||
_, t2 := s.createToken(t, ctx, u.ID, "b")
|
||||
_ = s.RevokeToken(ctx, u.ID, t1.ID)
|
||||
|
||||
list, err := s.ListTokens(ctx, u.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("list: %v", err)
|
||||
}
|
||||
if len(list) != 2 {
|
||||
t.Fatalf("expected 2 tokens (incl revoked), got %d", len(list))
|
||||
}
|
||||
for _, tk := range list {
|
||||
if tk.ID == t1.ID && tk.RevokedAt == nil {
|
||||
t.Fatal("revoked token should have RevokedAt set")
|
||||
}
|
||||
if tk.ID == t2.ID && tk.RevokedAt != nil {
|
||||
t.Fatal("active token should not be revoked")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestListTokens_ScopedToUser(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
alice := s.createUser(t, ctx, "alice", "p")
|
||||
bob := s.createUser(t, ctx, "bob", "p")
|
||||
s.createToken(t, ctx, alice.ID, "alice-token")
|
||||
s.createToken(t, ctx, bob.ID, "bob-token")
|
||||
|
||||
aliceList, _ := s.ListTokens(ctx, alice.ID)
|
||||
if len(aliceList) != 1 || aliceList[0].UserID != alice.ID {
|
||||
t.Fatalf("alice should see only her token, got %d", len(aliceList))
|
||||
}
|
||||
}
|
||||
|
||||
func TestTouchToken(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s := newTestStore(t)
|
||||
u := s.createUser(t, ctx, "alice", "p")
|
||||
plaintext, stored := s.createToken(t, ctx, u.ID, "laptop")
|
||||
|
||||
if stored.LastUsedAt != nil {
|
||||
t.Fatal("fresh token should have nil LastUsedAt")
|
||||
}
|
||||
if err := s.TouchToken(ctx, stored.ID); err != nil {
|
||||
t.Fatalf("touch: %v", err)
|
||||
}
|
||||
// Look up via the hash to confirm last_used_at was written.
|
||||
got, err := s.GetTokenByHash(ctx, crypto.HashToken(plaintext))
|
||||
if err != nil {
|
||||
t.Fatalf("get by hash: %v", err)
|
||||
}
|
||||
if got.LastUsedAt == nil {
|
||||
t.Fatal("expected LastUsedAt set after touch")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user