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:
2026-06-28 16:57:34 -05:00
commit 981587e83d
83 changed files with 12812 additions and 0 deletions
+19
View File
@@ -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
}
+80
View File
@@ -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
}
+53
View File
@@ -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
}
+102
View File
@@ -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)
}
}
+76
View File
@@ -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()
}
+236
View File
@@ -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,
)
}
+366
View File
@@ -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)
}
}
+125
View File
@@ -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
}
+253
View File
@@ -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])
}
}
+50
View File
@@ -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 }
+87
View File
@@ -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
}
+190
View File
@@ -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
}
+204
View File
@@ -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")
}
}
+105
View File
@@ -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
}
+127
View File
@@ -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
}
+299
View File
@@ -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")
}
}