Files
sub-store/internal/database/grant_repo.go
T

242 lines
6.4 KiB
Go

package database
import (
"database/sql"
"time"
"github.com/jmoiron/sqlx"
"github.com/google/uuid"
"github.com/peterqiu0516/sub-store/internal/model"
"github.com/peterqiu0516/sub-store/internal/util"
)
type GrantRepo struct {
db *sqlx.DB
}
func NewGrantRepo(db *sqlx.DB) *GrantRepo {
return &GrantRepo{db: db}
}
type grantRow struct {
ID string `db:"id"`
TokenHash string `db:"token_hash"`
ResourceType string `db:"resource_type"`
ResourceID string `db:"resource_id"`
Target string `db:"target"`
ExpiresAt *int64 `db:"expires_at"`
Enabled int `db:"enabled"`
CreatedAt int64 `db:"created_at"`
UpdatedAt int64 `db:"updated_at"`
}
// CreateGrant creates a new download grant, returning the grant record and the plaintext token.
func (r *GrantRepo) Create(resourceType, resourceID, target string, expiresAt *int64) (model.DownloadGrantRecord, string, error) {
now := time.Now().UnixMilli()
id := uuid.New().String()
token, err := util.RandomToken()
if err != nil {
return model.DownloadGrantRecord{}, "", err
}
tokenHash := util.SHA256Hex(token)
_, err = r.db.Exec(
`INSERT INTO download_grants (id, token_hash, resource_type, resource_id, target, expires_at, enabled, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?, 1, ?, ?)`,
id, tokenHash, resourceType, resourceID, target, expiresAt, now, now,
)
if err != nil {
return model.DownloadGrantRecord{}, "", err
}
rec := model.DownloadGrantRecord{
ID: id,
ResourceType: resourceType,
ResourceId: resourceID,
Target: target,
ExpiresAt: expiresAt,
Enabled: true,
CreatedAt: now,
UpdatedAt: now,
}
return rec, token, nil
}
func (r *GrantRepo) List() ([]model.DownloadGrantRecord, error) {
var rows []grantRow
if err := r.db.Select(&rows, "SELECT * FROM download_grants ORDER BY created_at DESC"); err != nil {
return nil, err
}
result := make([]model.DownloadGrantRecord, 0, len(rows))
for _, row := range rows {
result = append(result, grantFromRow(row))
}
return result, nil
}
func (r *GrantRepo) Get(id string) (*model.DownloadGrantRecord, error) {
var row grantRow
if err := r.db.Get(&row, "SELECT * FROM download_grants WHERE id = ?", id); err != nil {
if err == sql.ErrNoRows {
return nil, nil
}
return nil, err
}
rec := grantFromRow(row)
return &rec, nil
}
// GetSnapshot returns the grant record with the token hash for snapshot/restore.
func (r *GrantRepo) GetSnapshot(id string) (map[string]any, error) {
var row grantRow
if err := r.db.Get(&row, "SELECT * FROM download_grants WHERE id = ?", id); err != nil {
if err == sql.ErrNoRows {
return nil, nil
}
return nil, err
}
snapshot := map[string]any{
"id": row.ID,
"tokenHash": row.TokenHash,
"resourceType": row.ResourceType,
"resourceId": row.ResourceID,
"target": row.Target,
"expiresAt": row.ExpiresAt,
"enabled": row.Enabled != 0,
"createdAt": row.CreatedAt,
"updatedAt": row.UpdatedAt,
}
return snapshot, nil
}
func (r *GrantRepo) Update(id string, enabled *bool, expiresAt *int64) (*model.DownloadGrantRecord, error) {
existing, err := r.Get(id)
if err != nil {
return nil, err
}
if existing == nil {
return nil, nil
}
now := time.Now().UnixMilli()
if enabled != nil {
existing.Enabled = *enabled
}
if expiresAt != nil {
existing.ExpiresAt = expiresAt
}
// If expiresAt is explicitly set to 0, treat as nil (never expire)
if expiresAt != nil && *expiresAt == 0 {
existing.ExpiresAt = nil
}
_, err = r.db.Exec(
"UPDATE download_grants SET enabled = ?, expires_at = ?, updated_at = ? WHERE id = ?",
boolToInt(existing.Enabled), existing.ExpiresAt, now, id,
)
if err != nil {
return nil, err
}
return existing, nil
}
func (r *GrantRepo) Delete(id string) error {
_, err := r.db.Exec("DELETE FROM download_grants WHERE id = ?", id)
return err
}
// AuthorizeScoped checks if a token is valid for a scoped download.
func (r *GrantRepo) AuthorizeScoped(token, resourceType, resourceID, target string) bool {
if token == "" {
return false
}
tokenHash := util.SHA256Hex(token)
var row grantRow
err := r.db.Get(&row,
`SELECT * FROM download_grants WHERE token_hash = ? AND enabled = 1 AND resource_type = ? AND resource_id = ? LIMIT 1`,
tokenHash, resourceType, resourceID,
)
if err != nil {
return false
}
if row.ExpiresAt != nil && *row.ExpiresAt <= time.Now().UnixMilli() {
return false
}
restrictedTarget := model.NormalizeTargetAlias(row.Target)
return restrictedTarget == "" || restrictedTarget == target
}
// RestoreFromSnapshot inserts a grant from a recycled snapshot.
// Per review-resolution #38: restores tokenHash to download_grants table.
func (r *GrantRepo) RestoreFromSnapshot(snapshot map[string]any) error {
now := time.Now().UnixMilli()
id := getString(snapshot, "id")
tokenHash := getString(snapshot, "tokenHash")
resourceType := getString(snapshot, "resourceType")
resourceID := getString(snapshot, "resourceId")
if resourceType != "collection" {
resourceType = "source"
}
target := getString(snapshot, "target")
enabled := true
if e, ok := snapshot["enabled"].(bool); ok && !e {
enabled = false
}
var expiresAt *int64
if e, ok := snapshot["expiresAt"]; ok && e != nil {
if n, ok := e.(float64); ok && n > 0 {
v := int64(n)
expiresAt = &v
}
}
createdAt := getInt64(snapshot, "createdAt")
if createdAt == 0 {
createdAt = now
}
_, err := r.db.Exec(
`INSERT INTO download_grants (id, token_hash, resource_type, resource_id, target, expires_at, enabled, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`,
id, tokenHash, resourceType, resourceID, target, expiresAt, boolToInt(enabled), createdAt, now,
)
return err
}
func grantFromRow(row grantRow) model.DownloadGrantRecord {
resourceType := "source"
if row.ResourceType == "collection" {
resourceType = "collection"
}
return model.DownloadGrantRecord{
ID: row.ID,
ResourceType: resourceType,
ResourceId: row.ResourceID,
Target: row.Target,
ExpiresAt: row.ExpiresAt,
Enabled: row.Enabled != 0,
CreatedAt: row.CreatedAt,
UpdatedAt: row.UpdatedAt,
}
}
func getString(m map[string]any, key string) string {
if v, ok := m[key]; ok {
if s, ok := v.(string); ok {
return s
}
}
return ""
}
func getInt64(m map[string]any, key string) int64 {
if v, ok := m[key]; ok {
switch n := v.(type) {
case int64:
return n
case float64:
return int64(n)
}
}
return 0
}