fix server location on device registration

This commit is contained in:
m
2026-07-28 21:45:02 +02:00
parent 7af44152d7
commit 7cc425d28b
7 changed files with 158 additions and 31 deletions
+2 -2
View File
@@ -30,11 +30,11 @@ func main() {
queries := db.New(database) queries := db.New(database)
resourceSvc := service.NewResourceService(queries, cfg) resourceSvc := service.NewResourceService(database, queries, cfg)
ocrSvc := service.NewOCRService(cfg, resourceSvc) ocrSvc := service.NewOCRService(cfg, resourceSvc)
conversionSvc := service.NewConversionService(queries, cfg) conversionSvc := service.NewConversionService(queries, cfg)
urlSvc := service.NewURLService(cfg.HMACSecret, cfg.ServerHost, cfg.URLExpiryMinutes) 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 { if err != nil {
log.Fatalf("Failed to create auth service: %v", err) log.Fatalf("Failed to create auth service: %v", err)
} }
+46 -6
View File
@@ -4,6 +4,7 @@ import (
"context" "context"
"crypto/rand" "crypto/rand"
"crypto/sha256" "crypto/sha256"
"database/sql"
"encoding/base64" "encoding/base64"
"errors" "errors"
"fmt" "fmt"
@@ -30,6 +31,7 @@ const (
) )
type AuthService struct { type AuthService struct {
db *sql.DB
queries *db.Queries queries *db.Queries
key paseto.V4SymmetricKey key paseto.V4SymmetricKey
parser *paseto.Parser parser *paseto.Parser
@@ -45,7 +47,7 @@ type UserResponse struct {
Username string `json:"username"` 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) key, err := paseto.V4SymmetricKeyFromHex(cfg.PASETOKey)
if err != nil { if err != nil {
return nil, fmt.Errorf("invalid paseto key: %w", err) 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()) parser.AddRule(paseto.NotExpired())
return &AuthService{ return &AuthService{
db: database,
queries: queries, queries: queries,
key: key, key: key,
parser: &parser, 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) 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, Username: username,
PasswordHash: hash, 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) return nil, nil, fmt.Errorf("create user: %w", err)
} }
tokens, err := s.generateTokens(ctx, user.ID.String()) if _, err := qtx.CreateStorageLocation(ctx, db.CreateStorageLocationParams{
if err != nil { UserID: user.ID,
return nil, nil, fmt.Errorf("generate tokens: %w", err) 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) { 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';
+67 -22
View File
@@ -4,6 +4,7 @@ import (
"context" "context"
"database/sql" "database/sql"
"encoding/hex" "encoding/hex"
"errors"
"fmt" "fmt"
"mime/multipart" "mime/multipart"
"os" "os"
@@ -17,12 +18,13 @@ import (
) )
type ResourceService struct { type ResourceService struct {
db *sql.DB
queries *db.Queries queries *db.Queries
cfg *config.Config cfg *config.Config
} }
func NewResourceService(queries *db.Queries, cfg *config.Config) *ResourceService { func NewResourceService(database *sql.DB, queries *db.Queries, cfg *config.Config) *ResourceService {
return &ResourceService{queries: queries, cfg: cfg} return &ResourceService{db: database, queries: queries, cfg: cfg}
} }
func (s *ResourceService) Upload(file *multipart.FileHeader, ownerID string) (*model.UploadResult, error) { 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) 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, Checksum: checksum,
OwnerID: ownerUUID, OwnerID: ownerUUID,
}) })
if err == nil && existing.ID != uuid.Nil { if err == nil && existing.ID != uuid.Nil {
os.Remove(dst) os.Remove(dst)
placement, err := s.queries.GetServerPlacementByResource(context.Background(), existing.ID) placement, err := s.queries.GetServerPlacementByResource(ctx, existing.ID)
if err == nil { if err == nil {
return &model.UploadResult{ return &model.UploadResult{
ID: existing.ID.String(), ID: existing.ID.String(),
@@ -68,9 +72,18 @@ func (s *ResourceService) Upload(file *multipart.FileHeader, ownerID string) (*m
MimeType: existing.MimeType, MimeType: existing.MimeType,
}, nil }, 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, Name: file.Filename,
MimeType: file.Header.Get("Content-Type"), MimeType: file.Header.Get("Content-Type"),
Size: info.Size(), 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) 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 { if err != nil {
return nil, fmt.Errorf("create server placement: %w", err) 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) 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{ return &model.UploadResult{
ID: dbResource.ID.String(), ID: dbResource.ID.String(),
Name: dbResource.Name, 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) { 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 { 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, ResourceID: resourceID,
StorageLocationID: serverLoc.ID, StorageLocationID: serverLoc.ID,
Status: "synced", Status: "synced",
@@ -118,16 +154,6 @@ func (s *ResourceService) ensureServerPlacement(resourceID, ownerID uuid.UUID, d
return placement, nil 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) { func (s *ResourceService) List(ownerID string) ([]model.Resource, error) {
ownerUUID, _ := uuid.Parse(ownerID) ownerUUID, _ := uuid.Parse(ownerID)
dbResources, err := s.queries.ListResourcesByOwner(context.Background(), ownerUUID) 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) { func (s *ResourceService) CreateFolder(name, ownerID string) (*model.Resource, error) {
ownerUUID, _ := uuid.Parse(ownerID) 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, Name: name,
OwnerID: ownerUUID, OwnerID: ownerUUID,
}) })
@@ -251,10 +287,19 @@ func (s *ResourceService) CreateFolder(name, ownerID string) (*model.Resource, e
return nil, fmt.Errorf("create folder: %w", err) 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) 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) m := dbResourceToModel(r, nil)
return &m, nil return &m, nil
} }
+28 -1
View File
@@ -1,4 +1,5 @@
import { useEffect, useCallback, useRef } from 'react'; import { useEffect, useCallback, useRef } from 'react';
import { createMMKV } from 'react-native-mmkv';
import { useQueryClient } from '@tanstack/react-query'; import { useQueryClient } from '@tanstack/react-query';
import { File, UploadType } from 'expo-file-system'; import { File, UploadType } from 'expo-file-system';
import NetInfo, { NetInfoState } from '@react-native-community/netinfo'; import NetInfo, { NetInfoState } from '@react-native-community/netinfo';
@@ -9,6 +10,25 @@ import { API_BASE_URL, ENDPOINTS } from '../constants/api';
import { ApiError } from '../types'; import { ApiError } from '../types';
import { setIsSyncing } from './useSyncQueue'; import { setIsSyncing } from './useSyncQueue';
const MAX_RETRIES = 5;
const RETRY_MMKV_ID = 'vaultdrop-sync-retries';
const retryStorage = createMMKV({ id: RETRY_MMKV_ID });
function getRetryCount(fileId: string): number {
return retryStorage.getNumber(`${fileId}_retries`) ?? 0;
}
function incrementRetry(fileId: string): number {
const count = getRetryCount(fileId) + 1;
retryStorage.set(`${fileId}_retries`, count);
return count;
}
function resetRetry(fileId: string) {
retryStorage.remove(`${fileId}_retries`);
}
interface UploadResult { interface UploadResult {
name: string; name: string;
id: string; id: string;
@@ -100,13 +120,20 @@ export function useAutoSync() {
name: entry.name, name: entry.name,
}); });
resetRetry(entry.id);
fileStore.updatePartial(entry.id, { fileStore.updatePartial(entry.id, {
backendId: uploaded.id, backendId: uploaded.id,
syncStatus: 'synced', syncStatus: 'synced',
source: 'synced', source: 'synced',
}); });
} catch (err) { } catch (err) {
// upload failed, will retry on next cycle const retries = incrementRetry(entry.id);
if (retries >= MAX_RETRIES) {
fileStore.updatePartial(entry.id, {
syncStatus: 'error',
});
resetRetry(entry.id);
}
} }
} }
+12
View File
@@ -568,4 +568,16 @@ export const fileStore = {
.all() as FileRow[]; .all() as FileRow[];
return rows.map((r) => rowToRecord(r, getTagsForFile(r.id))); return rows.map((r) => rowToRecord(r, getTagsForFile(r.id)));
}, },
getErrorFiles(): FileRecord[] {
const d = getDb();
const rows = d.select().from(files)
.where(eq(files.syncStatus, 'error'))
.all() as FileRow[];
return rows.map((r) => rowToRecord(r, getTagsForFile(r.id)));
},
resetSyncError(id: string) {
this.updatePartial(id, { syncStatus: 'local' });
},
}; };