242 lines
6.4 KiB
Go
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
|
|
}
|