fix server location on device registration
This commit is contained in:
@@ -30,11 +30,11 @@ func main() {
|
||||
|
||||
queries := db.New(database)
|
||||
|
||||
resourceSvc := service.NewResourceService(queries, cfg)
|
||||
resourceSvc := service.NewResourceService(database, queries, cfg)
|
||||
ocrSvc := service.NewOCRService(cfg, resourceSvc)
|
||||
conversionSvc := service.NewConversionService(queries, cfg)
|
||||
urlSvc := service.NewURLService(cfg.HMACSecret, cfg.ServerHost, cfg.URLExpiryMinutes)
|
||||
authSvc, err := auth.NewAuthService(queries, cfg)
|
||||
authSvc, err := auth.NewAuthService(database, queries, cfg)
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to create auth service: %v", err)
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
@@ -30,6 +31,7 @@ const (
|
||||
)
|
||||
|
||||
type AuthService struct {
|
||||
db *sql.DB
|
||||
queries *db.Queries
|
||||
key paseto.V4SymmetricKey
|
||||
parser *paseto.Parser
|
||||
@@ -45,7 +47,7 @@ type UserResponse struct {
|
||||
Username string `json:"username"`
|
||||
}
|
||||
|
||||
func NewAuthService(queries *db.Queries, cfg *config.Config) (*AuthService, error) {
|
||||
func NewAuthService(database *sql.DB, queries *db.Queries, cfg *config.Config) (*AuthService, error) {
|
||||
key, err := paseto.V4SymmetricKeyFromHex(cfg.PASETOKey)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid paseto key: %w", err)
|
||||
@@ -55,6 +57,7 @@ func NewAuthService(queries *db.Queries, cfg *config.Config) (*AuthService, erro
|
||||
parser.AddRule(paseto.NotExpired())
|
||||
|
||||
return &AuthService{
|
||||
db: database,
|
||||
queries: queries,
|
||||
key: key,
|
||||
parser: &parser,
|
||||
@@ -80,7 +83,15 @@ func (s *AuthService) Register(ctx context.Context, username, password string) (
|
||||
return nil, nil, fmt.Errorf("hash password: %w", err)
|
||||
}
|
||||
|
||||
user, err := s.queries.CreateUser(ctx, db.CreateUserParams{
|
||||
tx, err := s.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("begin tx: %w", err)
|
||||
}
|
||||
defer tx.Rollback()
|
||||
|
||||
qtx := s.queries.WithTx(tx)
|
||||
|
||||
user, err := qtx.CreateUser(ctx, db.CreateUserParams{
|
||||
Username: username,
|
||||
PasswordHash: hash,
|
||||
})
|
||||
@@ -88,12 +99,41 @@ func (s *AuthService) Register(ctx context.Context, username, password string) (
|
||||
return nil, nil, fmt.Errorf("create user: %w", err)
|
||||
}
|
||||
|
||||
tokens, err := s.generateTokens(ctx, user.ID.String())
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("generate tokens: %w", err)
|
||||
if _, err := qtx.CreateStorageLocation(ctx, db.CreateStorageLocationParams{
|
||||
UserID: user.ID,
|
||||
DeviceName: "VaultDrop Server",
|
||||
Role: "server",
|
||||
}); err != nil {
|
||||
return nil, nil, fmt.Errorf("create server location: %w", err)
|
||||
}
|
||||
|
||||
return tokens, &UserResponse{ID: user.ID.String(), Username: user.Username}, nil
|
||||
accessToken, err := s.createAccessToken(user.ID.String())
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("create access token: %w", err)
|
||||
}
|
||||
|
||||
refreshToken, err := s.createRefreshToken(user.ID.String())
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("create refresh token: %w", err)
|
||||
}
|
||||
|
||||
refreshHash := hashToken(refreshToken)
|
||||
if _, err := qtx.CreateRefreshToken(ctx, db.CreateRefreshTokenParams{
|
||||
UserID: user.ID,
|
||||
TokenHash: refreshHash,
|
||||
ExpiresAt: time.Now().Add(refreshTokenTTL),
|
||||
}); err != nil {
|
||||
return nil, nil, fmt.Errorf("store refresh token: %w", err)
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, nil, fmt.Errorf("commit tx: %w", err)
|
||||
}
|
||||
|
||||
return &TokenPair{
|
||||
AccessToken: accessToken,
|
||||
RefreshToken: refreshToken,
|
||||
}, &UserResponse{ID: user.ID.String(), Username: user.Username}, nil
|
||||
}
|
||||
|
||||
func (s *AuthService) Login(ctx context.Context, username, password string) (*TokenPair, *UserResponse, error) {
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
DROP INDEX IF EXISTS idx_storage_locations_user_server;
|
||||
@@ -0,0 +1,2 @@
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_storage_locations_user_server
|
||||
ON storage_locations(user_id) WHERE role = 'server';
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"mime/multipart"
|
||||
"os"
|
||||
@@ -17,12 +18,13 @@ import (
|
||||
)
|
||||
|
||||
type ResourceService struct {
|
||||
db *sql.DB
|
||||
queries *db.Queries
|
||||
cfg *config.Config
|
||||
}
|
||||
|
||||
func NewResourceService(queries *db.Queries, cfg *config.Config) *ResourceService {
|
||||
return &ResourceService{queries: queries, cfg: cfg}
|
||||
func NewResourceService(database *sql.DB, queries *db.Queries, cfg *config.Config) *ResourceService {
|
||||
return &ResourceService{db: database, queries: queries, cfg: cfg}
|
||||
}
|
||||
|
||||
func (s *ResourceService) Upload(file *multipart.FileHeader, ownerID string) (*model.UploadResult, error) {
|
||||
@@ -53,13 +55,15 @@ func (s *ResourceService) Upload(file *multipart.FileHeader, ownerID string) (*m
|
||||
return nil, fmt.Errorf("parse owner id: %w", err)
|
||||
}
|
||||
|
||||
existing, err := s.queries.FindDuplicateByChecksum(context.Background(), db.FindDuplicateByChecksumParams{
|
||||
ctx := context.Background()
|
||||
|
||||
existing, err := s.queries.FindDuplicateByChecksum(ctx, db.FindDuplicateByChecksumParams{
|
||||
Checksum: checksum,
|
||||
OwnerID: ownerUUID,
|
||||
})
|
||||
if err == nil && existing.ID != uuid.Nil {
|
||||
os.Remove(dst)
|
||||
placement, err := s.queries.GetServerPlacementByResource(context.Background(), existing.ID)
|
||||
placement, err := s.queries.GetServerPlacementByResource(ctx, existing.ID)
|
||||
if err == nil {
|
||||
return &model.UploadResult{
|
||||
ID: existing.ID.String(),
|
||||
@@ -68,9 +72,18 @@ func (s *ResourceService) Upload(file *multipart.FileHeader, ownerID string) (*m
|
||||
MimeType: existing.MimeType,
|
||||
}, nil
|
||||
}
|
||||
return nil, fmt.Errorf("duplicate resource %s has no server placement — upload cannot proceed until resolved", existing.ID.String())
|
||||
}
|
||||
|
||||
dbResource, err := s.queries.CreateResource(context.Background(), db.CreateResourceParams{
|
||||
tx, err := s.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("begin tx: %w", err)
|
||||
}
|
||||
defer tx.Rollback()
|
||||
|
||||
qtx := s.queries.WithTx(tx)
|
||||
|
||||
dbResource, err := qtx.CreateResource(ctx, db.CreateResourceParams{
|
||||
Name: file.Filename,
|
||||
MimeType: file.Header.Get("Content-Type"),
|
||||
Size: info.Size(),
|
||||
@@ -81,15 +94,24 @@ func (s *ResourceService) Upload(file *multipart.FileHeader, ownerID string) (*m
|
||||
return nil, fmt.Errorf("create resource in db: %w", err)
|
||||
}
|
||||
|
||||
placement, err := s.ensureServerPlacement(dbResource.ID, ownerUUID, dst)
|
||||
placement, err := s.ensureServerPlacementQtx(qtx, dbResource.ID, ownerUUID, dst)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create server placement: %w", err)
|
||||
}
|
||||
|
||||
if err := s.ensureOwnerRebac(dbResource.ID, ownerUUID); err != nil {
|
||||
if _, err := qtx.CreateRebacRelation(ctx, db.CreateRebacRelationParams{
|
||||
ResourceID: dbResource.ID,
|
||||
SubjectUserID: ownerUUID,
|
||||
Role: "owner",
|
||||
GrantedBy: ownerUUID,
|
||||
}); err != nil {
|
||||
return nil, fmt.Errorf("create owner rebac: %w", err)
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, fmt.Errorf("commit tx: %w", err)
|
||||
}
|
||||
|
||||
return &model.UploadResult{
|
||||
ID: dbResource.ID.String(),
|
||||
Name: dbResource.Name,
|
||||
@@ -99,12 +121,26 @@ func (s *ResourceService) Upload(file *multipart.FileHeader, ownerID string) (*m
|
||||
}
|
||||
|
||||
func (s *ResourceService) ensureServerPlacement(resourceID, ownerID uuid.UUID, dst string) (db.ResourcePlacement, error) {
|
||||
serverLoc, err := s.queries.GetServerStorageLocation(context.Background(), ownerID)
|
||||
return s.ensureServerPlacementQtx(s.queries, resourceID, ownerID, dst)
|
||||
}
|
||||
|
||||
func (s *ResourceService) ensureServerPlacementQtx(q *db.Queries, resourceID, ownerID uuid.UUID, dst string) (db.ResourcePlacement, error) {
|
||||
ctx := context.Background()
|
||||
serverLoc, err := q.GetServerStorageLocation(ctx, ownerID)
|
||||
if err != nil {
|
||||
return db.ResourcePlacement{}, fmt.Errorf("get server location: %w", err)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
serverLoc, err = q.CreateStorageLocation(ctx, db.CreateStorageLocationParams{
|
||||
UserID: ownerID,
|
||||
DeviceName: "VaultDrop Server",
|
||||
Role: "server",
|
||||
})
|
||||
}
|
||||
if err != nil {
|
||||
return db.ResourcePlacement{}, fmt.Errorf("get/create server location: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
placement, err := s.queries.CreatePlacement(context.Background(), db.CreatePlacementParams{
|
||||
placement, err := q.CreatePlacement(ctx, db.CreatePlacementParams{
|
||||
ResourceID: resourceID,
|
||||
StorageLocationID: serverLoc.ID,
|
||||
Status: "synced",
|
||||
@@ -118,16 +154,6 @@ func (s *ResourceService) ensureServerPlacement(resourceID, ownerID uuid.UUID, d
|
||||
return placement, nil
|
||||
}
|
||||
|
||||
func (s *ResourceService) ensureOwnerRebac(resourceID, ownerID uuid.UUID) error {
|
||||
_, err := s.queries.CreateRebacRelation(context.Background(), db.CreateRebacRelationParams{
|
||||
ResourceID: resourceID,
|
||||
SubjectUserID: ownerID,
|
||||
Role: "owner",
|
||||
GrantedBy: ownerID,
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *ResourceService) List(ownerID string) ([]model.Resource, error) {
|
||||
ownerUUID, _ := uuid.Parse(ownerID)
|
||||
dbResources, err := s.queries.ListResourcesByOwner(context.Background(), ownerUUID)
|
||||
@@ -243,7 +269,17 @@ func (s *ResourceService) MoveResources(resourceIDs []string, parentResourceID *
|
||||
|
||||
func (s *ResourceService) CreateFolder(name, ownerID string) (*model.Resource, error) {
|
||||
ownerUUID, _ := uuid.Parse(ownerID)
|
||||
r, err := s.queries.CreateFolder(context.Background(), db.CreateFolderParams{
|
||||
ctx := context.Background()
|
||||
|
||||
tx, err := s.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("begin tx: %w", err)
|
||||
}
|
||||
defer tx.Rollback()
|
||||
|
||||
qtx := s.queries.WithTx(tx)
|
||||
|
||||
r, err := qtx.CreateFolder(ctx, db.CreateFolderParams{
|
||||
Name: name,
|
||||
OwnerID: ownerUUID,
|
||||
})
|
||||
@@ -251,10 +287,19 @@ func (s *ResourceService) CreateFolder(name, ownerID string) (*model.Resource, e
|
||||
return nil, fmt.Errorf("create folder: %w", err)
|
||||
}
|
||||
|
||||
if err := s.ensureOwnerRebac(r.ID, ownerUUID); err != nil {
|
||||
if _, err := qtx.CreateRebacRelation(ctx, db.CreateRebacRelationParams{
|
||||
ResourceID: r.ID,
|
||||
SubjectUserID: ownerUUID,
|
||||
Role: "owner",
|
||||
GrantedBy: ownerUUID,
|
||||
}); err != nil {
|
||||
return nil, fmt.Errorf("create owner rebac: %w", err)
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, fmt.Errorf("commit tx: %w", err)
|
||||
}
|
||||
|
||||
m := dbResourceToModel(r, nil)
|
||||
return &m, nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user