From 09a22900cdbd54eaddeae87a0533de7fc79e519a Mon Sep 17 00:00:00 2001 From: Tony Date: Mon, 15 Jun 2026 15:17:02 +0800 Subject: [PATCH] feat: implement app access control and upload token validation in storage --- cmd/server/main.go | 1 + internal/api/handlers/access.go | 46 +++++++++++++++ internal/api/handlers/admin.go | 12 ++++ internal/api/handlers/apps.go | 13 ++-- internal/api/handlers/device.go | 13 +++- internal/api/handlers/patches.go | 75 ++++++++++++++++++++++++ internal/api/handlers/releases.go | 59 +++++++++++++++++++ internal/api/handlers/router.go | 2 +- internal/api/handlers/storage_handler.go | 13 ++++ internal/storage/factory.go | 19 +++--- internal/storage/local.go | 66 +++++++++++++++++++-- 11 files changed, 297 insertions(+), 22 deletions(-) create mode 100644 internal/api/handlers/access.go diff --git a/cmd/server/main.go b/cmd/server/main.go index 3233505..7cfaf89 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -64,6 +64,7 @@ func main() { Type: cfg.Storage.Driver, LocalDir: cfg.Storage.LocalDir, ServerBaseURL: cfg.Server.BaseURL, + UploadSecret: cfg.Auth.JWTSecret, S3Endpoint: cfg.Storage.S3Endpoint, S3AccessKey: cfg.Storage.S3AccessKey, S3SecretKey: cfg.Storage.S3SecretKey, diff --git a/internal/api/handlers/access.go b/internal/api/handlers/access.go new file mode 100644 index 0000000..79f1b1d --- /dev/null +++ b/internal/api/handlers/access.go @@ -0,0 +1,46 @@ +package handlers + +import ( + "net/http" + + "github.com/google/uuid" + "github.com/shorebird-server/internal/api/middleware" + "github.com/shorebird-server/internal/db" +) + +func requireAppAccess(w http.ResponseWriter, r *http.Request, store db.Store, appID uuid.UUID, manage bool) (*db.AppRow, bool) { + claims := middleware.GetClaims(r) + if claims == nil { + respondError(w, http.StatusUnauthorized, "Unauthorized", nil) + return nil, false + } + app, err := store.GetAppByID(r.Context(), appID) + if err != nil { + respondError(w, http.StatusNotFound, "App not found", nil) + return nil, false + } + if canAccessOrganization(r, store, claims.UserID, app.OrganizationID, manage) { + return app, true + } + if manage { + respondError(w, http.StatusForbidden, "Organization admin access required", nil) + } else { + respondError(w, http.StatusForbidden, "Organization access required", nil) + } + return nil, false +} + +func canAccessOrganization(r *http.Request, store db.Store, userID, orgID int, manage bool) bool { + user, _ := store.GetUserByID(r.Context(), userID) + if user != nil && user.IsAdmin { + return true + } + role, err := store.GetOrganizationRole(r.Context(), userID, orgID) + if err != nil { + return false + } + if manage { + return role == "admin" + } + return role != "" +} diff --git a/internal/api/handlers/admin.go b/internal/api/handlers/admin.go index daece12..9d593b8 100644 --- a/internal/api/handlers/admin.go +++ b/internal/api/handlers/admin.go @@ -17,6 +17,9 @@ type AdminHandler struct { // AddTargetDevice handles POST /api/v1/admin/patches/{patchId}/target-devices // Restricts a patch to only be delivered to specific devices (by client_id). func (h *AdminHandler) AddTargetDevice(w http.ResponseWriter, r *http.Request) { + if !h.requireAdmin(w, r) { + return + } patchID, err := strconv.Atoi(chi.URLParam(r, "patchId")) if err != nil { respondError(w, http.StatusBadRequest, "Invalid patch ID", nil) @@ -41,6 +44,9 @@ func (h *AdminHandler) AddTargetDevice(w http.ResponseWriter, r *http.Request) { // RemoveTargetDevice handles DELETE /api/v1/admin/patches/{patchId}/target-devices/{clientId} func (h *AdminHandler) RemoveTargetDevice(w http.ResponseWriter, r *http.Request) { + if !h.requireAdmin(w, r) { + return + } // For simplicity, target device removal is handled by the database cascade // In a full implementation, add a RemovePatchTargetDevice DB method respondError(w, http.StatusNotImplemented, "Not yet implemented", nil) @@ -48,6 +54,9 @@ func (h *AdminHandler) RemoveTargetDevice(w http.ResponseWriter, r *http.Request // GetTargetDevices handles GET /api/v1/admin/patches/{patchId}/target-devices func (h *AdminHandler) GetTargetDevices(w http.ResponseWriter, r *http.Request) { + if !h.requireAdmin(w, r) { + return + } patchID, err := strconv.Atoi(chi.URLParam(r, "patchId")) if err != nil { respondError(w, http.StatusBadRequest, "Invalid patch ID", nil) @@ -69,6 +78,9 @@ func (h *AdminHandler) GetTargetDevices(w http.ResponseWriter, r *http.Request) // GetPatchEvents handles GET /api/v1/admin/apps/{appId}/events // Returns patch events for analytics. func (h *AdminHandler) GetPatchEvents(w http.ResponseWriter, r *http.Request) { + if !h.requireAdmin(w, r) { + return + } // This would query patch_events table for the app // For now, return a stub respondJSON(w, http.StatusOK, map[string]interface{}{ diff --git a/internal/api/handlers/apps.go b/internal/api/handlers/apps.go index abf49ba..cb31d9a 100644 --- a/internal/api/handlers/apps.go +++ b/internal/api/handlers/apps.go @@ -285,12 +285,7 @@ func (h *AppHandler) TransferApp(w http.ResponseWriter, r *http.Request) { } func (h *AppHandler) canManageOrganization(r *http.Request, userID, orgID int) bool { - user, _ := h.DB.GetUserByID(r.Context(), userID) - if user != nil && user.IsAdmin { - return true - } - role, err := h.DB.GetOrganizationRole(r.Context(), userID, orgID) - return err == nil && role == "admin" + return canAccessOrganization(r, h.DB, userID, orgID, true) } // ChannelHandler handles channel-related API endpoints. @@ -306,6 +301,9 @@ func (h *ChannelHandler) GetChannels(w http.ResponseWriter, r *http.Request) { respondError(w, http.StatusBadRequest, "Invalid app ID", nil) return } + if _, ok := requireAppAccess(w, r, h.DB, id, false); !ok { + return + } channels, err := h.DB.GetChannelsByAppID(r.Context(), id) if err != nil { @@ -333,6 +331,9 @@ func (h *ChannelHandler) CreateChannel(w http.ResponseWriter, r *http.Request) { respondError(w, http.StatusBadRequest, "Invalid app ID", nil) return } + if _, ok := requireAppAccess(w, r, h.DB, id, true); !ok { + return + } var req models.CreateChannelRequest if err := decodeJSON(r, &req); err != nil { diff --git a/internal/api/handlers/device.go b/internal/api/handlers/device.go index 0ca54d1..17ff98e 100644 --- a/internal/api/handlers/device.go +++ b/internal/api/handlers/device.go @@ -4,6 +4,7 @@ import ( "net/http" "github.com/google/uuid" + "github.com/shorebird-server/internal/api/middleware" "github.com/shorebird-server/internal/db" "github.com/shorebird-server/internal/models" "github.com/shorebird-server/internal/storage" @@ -179,7 +180,17 @@ func (h *DeviceHandler) PatchEvents(w http.ResponseWriter, r *http.Request) { // RollbackPatch handles POST /api/v1/apps/{appId}/patches/rollback (admin endpoint) func (h *DeviceHandler) RollbackPatch(w http.ResponseWriter, r *http.Request) { - // This is an admin operation. For now, allow authenticated users. + claims := middleware.GetClaims(r) + if claims == nil { + respondError(w, http.StatusUnauthorized, "Unauthorized", nil) + return + } + user, err := h.DB.GetUserByID(r.Context(), claims.UserID) + if err != nil || !user.IsAdmin { + respondError(w, http.StatusForbidden, "Admin access required", nil) + return + } + var req struct { ReleaseID int `json:"release_id"` PatchNumber int `json:"patch_number"` diff --git a/internal/api/handlers/patches.go b/internal/api/handlers/patches.go index f3272dd..54e3faf 100644 --- a/internal/api/handlers/patches.go +++ b/internal/api/handlers/patches.go @@ -7,6 +7,7 @@ import ( "time" "github.com/go-chi/chi/v5" + "github.com/google/uuid" "github.com/shorebird-server/internal/db" "github.com/shorebird-server/internal/models" "github.com/shorebird-server/internal/storage" @@ -21,6 +22,14 @@ type PatchHandler struct { // CreatePatch handles POST /api/v1/apps/{appId}/patches func (h *PatchHandler) CreatePatch(w http.ResponseWriter, r *http.Request) { appID := chi.URLParam(r, "appId") + parsedAppID, err := uuid.Parse(appID) + if err != nil { + respondError(w, http.StatusBadRequest, "Invalid app ID", nil) + return + } + if _, ok := requireAppAccess(w, r, h.DB, parsedAppID, true); !ok { + return + } var req models.CreatePatchRequest if err := decodeJSON(r, &req); err != nil { @@ -62,6 +71,14 @@ func (h *PatchHandler) CreatePatch(w http.ResponseWriter, r *http.Request) { // CreatePatchArtifact handles POST /api/v1/apps/{appId}/patches/{patchId}/artifacts func (h *PatchHandler) CreatePatchArtifact(w http.ResponseWriter, r *http.Request) { appID := chi.URLParam(r, "appId") + parsedAppID, err := uuid.Parse(appID) + if err != nil { + respondError(w, http.StatusBadRequest, "Invalid app ID", nil) + return + } + if _, ok := requireAppAccess(w, r, h.DB, parsedAppID, true); !ok { + return + } patchID, err := strconv.Atoi(chi.URLParam(r, "patchId")) if err != nil { respondError(w, http.StatusBadRequest, "Invalid patch ID", nil) @@ -95,6 +112,10 @@ func (h *PatchHandler) CreatePatchArtifact(w http.ResponseWriter, r *http.Reques respondError(w, http.StatusNotFound, "Release not found", nil) return } + if release.AppID != parsedAppID { + respondError(w, http.StatusNotFound, "Patch not found for this app", nil) + return + } // Generate a storage key for the public bucket filename := "dlc.vmcode" // standard patch artifact name @@ -136,11 +157,24 @@ func (h *PatchHandler) CreatePatchArtifact(w http.ResponseWriter, r *http.Reques // GetPatches handles GET /api/v1/apps/{appId}/releases/{releaseId}/patches func (h *PatchHandler) GetPatches(w http.ResponseWriter, r *http.Request) { + appID, err := uuid.Parse(chi.URLParam(r, "appId")) + if err != nil { + respondError(w, http.StatusBadRequest, "Invalid app ID", nil) + return + } + if _, ok := requireAppAccess(w, r, h.DB, appID, false); !ok { + return + } releaseID, err := strconv.Atoi(chi.URLParam(r, "releaseId")) if err != nil { respondError(w, http.StatusBadRequest, "Invalid release ID", nil) return } + release, err := h.DB.GetReleaseByID(r.Context(), releaseID) + if err != nil || release.AppID != appID { + respondError(w, http.StatusNotFound, "Release not found for this app", nil) + return + } patches, err := h.DB.GetPatchesByReleaseID(r.Context(), releaseID) if err != nil { @@ -172,11 +206,22 @@ func (h *PatchHandler) GetPatches(w http.ResponseWriter, r *http.Request) { // PromotePatch handles POST /api/v1/apps/{appId}/patches/promote func (h *PatchHandler) PromotePatch(w http.ResponseWriter, r *http.Request) { + appID, err := uuid.Parse(chi.URLParam(r, "appId")) + if err != nil { + respondError(w, http.StatusBadRequest, "Invalid app ID", nil) + return + } + if _, ok := requireAppAccess(w, r, h.DB, appID, true); !ok { + return + } var req models.PromotePatchRequest if err := decodeJSON(r, &req); err != nil { respondError(w, http.StatusBadRequest, "Invalid request body", nil) return } + if !h.patchBelongsToApp(w, r, req.PatchID, appID) { + return + } if err := h.DB.PromotePatch(r.Context(), req.PatchID, req.ChannelID); err != nil { respondError(w, http.StatusInternalServerError, "Failed to promote patch", nil) @@ -188,6 +233,22 @@ func (h *PatchHandler) PromotePatch(w http.ResponseWriter, r *http.Request) { // UpdatePatch handles PATCH /api/v1/apps/{appId}/patches/{patchId} func (h *PatchHandler) UpdatePatch(w http.ResponseWriter, r *http.Request) { + appID, err := uuid.Parse(chi.URLParam(r, "appId")) + if err != nil { + respondError(w, http.StatusBadRequest, "Invalid app ID", nil) + return + } + if _, ok := requireAppAccess(w, r, h.DB, appID, true); !ok { + return + } + patchID, err := strconv.Atoi(chi.URLParam(r, "patchId")) + if err != nil { + respondError(w, http.StatusBadRequest, "Invalid patch ID", nil) + return + } + if !h.patchBelongsToApp(w, r, patchID, appID) { + return + } var req models.UpdatePatchRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { respondError(w, http.StatusBadRequest, "Invalid request body", nil) @@ -198,3 +259,17 @@ func (h *PatchHandler) UpdatePatch(w http.ResponseWriter, r *http.Request) { // This would require an additional DB method respondJSON(w, http.StatusOK, map[string]string{"status": "ok"}) } + +func (h *PatchHandler) patchBelongsToApp(w http.ResponseWriter, r *http.Request, patchID int, appID uuid.UUID) bool { + patch, err := h.DB.GetPatchByID(r.Context(), patchID) + if err != nil { + respondError(w, http.StatusNotFound, "Patch not found", nil) + return false + } + release, err := h.DB.GetReleaseByID(r.Context(), patch.ReleaseID) + if err != nil || release.AppID != appID { + respondError(w, http.StatusNotFound, "Patch not found for this app", nil) + return false + } + return true +} diff --git a/internal/api/handlers/releases.go b/internal/api/handlers/releases.go index f798850..3c0c361 100644 --- a/internal/api/handlers/releases.go +++ b/internal/api/handlers/releases.go @@ -38,6 +38,9 @@ func (h *ReleaseHandler) GetReleases(w http.ResponseWriter, r *http.Request) { respondError(w, http.StatusBadRequest, "Invalid app ID", nil) return } + if _, ok := requireAppAccess(w, r, h.DB, appID, false); !ok { + return + } sideloadable := r.URL.Query().Get("sideloadable") == "true" @@ -81,6 +84,9 @@ func (h *ReleaseHandler) CreateRelease(w http.ResponseWriter, r *http.Request) { respondError(w, http.StatusBadRequest, "Invalid app ID", nil) return } + if _, ok := requireAppAccess(w, r, h.DB, appID, true); !ok { + return + } var req models.CreateReleaseRequest if err := decodeJSON(r, &req); err != nil { @@ -112,11 +118,22 @@ func (h *ReleaseHandler) CreateRelease(w http.ResponseWriter, r *http.Request) { // UpdateRelease handles PATCH /api/v1/apps/{appId}/releases/{releaseId} func (h *ReleaseHandler) UpdateRelease(w http.ResponseWriter, r *http.Request) { + appID, err := parseUUID(chi.URLParam(r, "appId")) + if err != nil { + respondError(w, http.StatusBadRequest, "Invalid app ID", nil) + return + } + if _, ok := requireAppAccess(w, r, h.DB, appID, true); !ok { + return + } releaseID, err := parseIntParam(r, "releaseId") if err != nil { respondError(w, http.StatusBadRequest, "Invalid release ID", nil) return } + if !h.releaseBelongsToApp(w, r, releaseID, appID) { + return + } var req models.UpdateReleaseRequest if err := decodeJSON(r, &req); err != nil { @@ -133,11 +150,22 @@ func (h *ReleaseHandler) UpdateRelease(w http.ResponseWriter, r *http.Request) { } func (h *ReleaseHandler) DeleteRelease(w http.ResponseWriter, r *http.Request) { + appID, err := parseUUID(chi.URLParam(r, "appId")) + if err != nil { + respondError(w, http.StatusBadRequest, "Invalid app ID", nil) + return + } + if _, ok := requireAppAccess(w, r, h.DB, appID, true); !ok { + return + } releaseID, err := parseIntParam(r, "releaseId") if err != nil { respondError(w, http.StatusBadRequest, "Invalid release ID", nil) return } + if !h.releaseBelongsToApp(w, r, releaseID, appID) { + return + } if err := h.DB.DeleteRelease(r.Context(), releaseID); err != nil { respondError(w, http.StatusInternalServerError, "Failed to delete release", strPtr(err.Error())) return @@ -148,11 +176,22 @@ func (h *ReleaseHandler) DeleteRelease(w http.ResponseWriter, r *http.Request) { // CreateReleaseArtifact handles POST /api/v1/apps/{appId}/releases/{releaseId}/artifacts func (h *ReleaseHandler) CreateReleaseArtifact(w http.ResponseWriter, r *http.Request) { appID := chi.URLParam(r, "appId") + parsedAppID, err := parseUUID(appID) + if err != nil { + respondError(w, http.StatusBadRequest, "Invalid app ID", nil) + return + } + if _, ok := requireAppAccess(w, r, h.DB, parsedAppID, true); !ok { + return + } releaseID, err := parseIntParam(r, "releaseId") if err != nil { respondError(w, http.StatusBadRequest, "Invalid release ID", nil) return } + if !h.releaseBelongsToApp(w, r, releaseID, parsedAppID) { + return + } // The CLI sends multipart form data. Parse it. if err := r.ParseMultipartForm(100 << 20); err != nil { // 100MB max @@ -216,11 +255,22 @@ func (h *ReleaseHandler) CreateReleaseArtifact(w http.ResponseWriter, r *http.Re // GetReleaseArtifacts handles GET /api/v1/apps/{appId}/releases/{releaseId}/artifacts func (h *ReleaseHandler) GetReleaseArtifacts(w http.ResponseWriter, r *http.Request) { + appID, err := parseUUID(chi.URLParam(r, "appId")) + if err != nil { + respondError(w, http.StatusBadRequest, "Invalid app ID", nil) + return + } + if _, ok := requireAppAccess(w, r, h.DB, appID, false); !ok { + return + } releaseID, err := parseIntParam(r, "releaseId") if err != nil { respondError(w, http.StatusBadRequest, "Invalid release ID", nil) return } + if !h.releaseBelongsToApp(w, r, releaseID, appID) { + return + } var archPtr, platformPtr *string if a := r.URL.Query().Get("arch"); a != "" { @@ -260,3 +310,12 @@ func (h *ReleaseHandler) GetReleaseArtifacts(w http.ResponseWriter, r *http.Requ respondJSON(w, http.StatusOK, map[string]interface{}{"artifacts": result}) } + +func (h *ReleaseHandler) releaseBelongsToApp(w http.ResponseWriter, r *http.Request, releaseID int, appID uuid.UUID) bool { + release, err := h.DB.GetReleaseByID(r.Context(), releaseID) + if err != nil || release.AppID != appID { + respondError(w, http.StatusNotFound, "Release not found for this app", nil) + return false + } + return true +} diff --git a/internal/api/handlers/router.go b/internal/api/handlers/router.go index ef5e3ba..bf80f0a 100644 --- a/internal/api/handlers/router.go +++ b/internal/api/handlers/router.go @@ -27,7 +27,7 @@ func NewRouter(authService *authpkg.Service, database db.Store, store storage.St AllowedMethods: []string{"GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"}, AllowedHeaders: []string{"Accept", "Authorization", "Content-Type", "X-Version"}, ExposedHeaders: []string{"Link"}, - AllowCredentials: true, + AllowCredentials: false, MaxAge: 300, })) diff --git a/internal/api/handlers/storage_handler.go b/internal/api/handlers/storage_handler.go index 4a08c05..ea12436 100644 --- a/internal/api/handlers/storage_handler.go +++ b/internal/api/handlers/storage_handler.go @@ -3,12 +3,15 @@ package handlers import ( "io" "net/http" + "strconv" "strings" "github.com/go-chi/chi/v5" "github.com/shorebird-server/internal/storage" ) +const maxStorageUploadBytes = 2 << 30 + // StorageHandler provides upload/download endpoints for local storage. type StorageHandler struct { store storage.Store @@ -21,8 +24,18 @@ func NewStorageHandler(store storage.Store) *StorageHandler { // Upload handles POST /storage/upload/{scope}/{key} func (h *StorageHandler) Upload(w http.ResponseWriter, r *http.Request) { + r.Body = http.MaxBytesReader(w, r.Body, maxStorageUploadBytes) scope := chi.URLParam(r, "scope") key := chi.URLParam(r, "*") + if v, ok := h.store.(interface { + ValidateUploadToken(scope, objectKey, token string, expires int64) bool + }); ok { + expires, _ := strconv.ParseInt(r.URL.Query().Get("expires"), 10, 64) + if !v.ValidateUploadToken(scope, key, r.URL.Query().Get("token"), expires) { + respondError(w, http.StatusForbidden, "Invalid or expired upload URL", nil) + return + } + } isPublic := scope == "patches" ct := r.Header.Get("Content-Type") diff --git a/internal/storage/factory.go b/internal/storage/factory.go index ab0ab88..36a0d44 100644 --- a/internal/storage/factory.go +++ b/internal/storage/factory.go @@ -8,16 +8,17 @@ type Config struct { Type string // Local settings - LocalDir string - ServerBaseURL string + LocalDir string + ServerBaseURL string + UploadSecret string // S3 settings - S3Endpoint string - S3AccessKey string - S3SecretKey string - S3UseSSL bool - S3ReleaseBucket string - S3PatchBucket string + S3Endpoint string + S3AccessKey string + S3SecretKey string + S3UseSSL bool + S3ReleaseBucket string + S3PatchBucket string } // NewStore creates the appropriate Store backend based on Config.Type. @@ -32,6 +33,6 @@ func NewStore(cfg Config) (Store, error) { if cfg.Type != "" && cfg.Type != "local" { fmt.Printf("storage: unknown type %q, falling back to local\n", cfg.Type) } - return NewLocalStore(cfg.LocalDir, cfg.ServerBaseURL) + return NewLocalStore(cfg.LocalDir, cfg.ServerBaseURL, cfg.UploadSecret) } } diff --git a/internal/storage/local.go b/internal/storage/local.go index ccf5fbc..dee6eca 100644 --- a/internal/storage/local.go +++ b/internal/storage/local.go @@ -2,10 +2,15 @@ package storage import ( "context" + "crypto/hmac" + "crypto/sha256" + "encoding/base64" "fmt" "io" "os" "path/filepath" + "strconv" + "strings" "time" ) @@ -18,11 +23,12 @@ type LocalStore struct { // serverBaseURL is the external URL of this server (e.g. http://localhost:8080). // Used when generating download URLs for device-side patch checks. serverBaseURL string + uploadSecret string } // NewLocalStore creates a local-filesystem storage backend. // baseDir is created automatically if it does not exist. -func NewLocalStore(baseDir, serverBaseURL string) (*LocalStore, error) { +func NewLocalStore(baseDir, serverBaseURL, uploadSecret string) (*LocalStore, error) { abs, err := filepath.Abs(baseDir) if err != nil { return nil, fmt.Errorf("local storage: %w", err) @@ -36,7 +42,7 @@ func NewLocalStore(baseDir, serverBaseURL string) (*LocalStore, error) { if serverBaseURL == "" { serverBaseURL = "http://localhost:8080" } - return &LocalStore{baseDir: abs, serverBaseURL: serverBaseURL}, nil + return &LocalStore{baseDir: abs, serverBaseURL: serverBaseURL, uploadSecret: uploadSecret}, nil } func (s *LocalStore) BackendName() string { return "Local filesystem" } @@ -49,7 +55,12 @@ func (s *LocalStore) GeneratePresignedUploadURL(_ context.Context, objectKey str if isPublic { scope = "patches" } - return fmt.Sprintf("%s/storage/upload/%s/%s", s.serverBaseURL, scope, objectKey), nil + base := fmt.Sprintf("%s/storage/upload/%s/%s", s.serverBaseURL, scope, objectKey) + if s.uploadSecret == "" { + return base, nil + } + expires := time.Now().Add(30 * time.Minute).Unix() + return fmt.Sprintf("%s?expires=%d&token=%s", base, expires, s.signUpload(scope, objectKey, expires)), nil } // GeneratePublicDownloadURL returns a public download URL for a patch. @@ -63,7 +74,10 @@ func (s *LocalStore) UploadObject(_ context.Context, objectKey string, reader io if isPublic { scope = "patches" } - target := filepath.Join(s.baseDir, scope, objectKey) + target, err := s.objectPath(scope, objectKey) + if err != nil { + return err + } if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { return fmt.Errorf("local storage: %w", err) } @@ -84,8 +98,50 @@ func (s *LocalStore) GetObject(_ context.Context, objectKey string, isPublic boo if isPublic { scope = "patches" } - return os.Open(filepath.Join(s.baseDir, scope, objectKey)) + target, err := s.objectPath(scope, objectKey) + if err != nil { + return nil, err + } + return os.Open(target) } // BaseDir returns the absolute base directory. func (s *LocalStore) BaseDir() string { return s.baseDir } + +func (s *LocalStore) ValidateUploadToken(scope, objectKey, token string, expires int64) bool { + if s.uploadSecret == "" { + return true + } + if token == "" || expires <= time.Now().Unix() { + return false + } + expected := s.signUpload(scope, objectKey, expires) + return hmac.Equal([]byte(token), []byte(expected)) +} + +func (s *LocalStore) signUpload(scope, objectKey string, expires int64) string { + mac := hmac.New(sha256.New, []byte(s.uploadSecret)) + io.WriteString(mac, scope) + io.WriteString(mac, "\n") + io.WriteString(mac, objectKey) + io.WriteString(mac, "\n") + io.WriteString(mac, strconv.FormatInt(expires, 10)) + return base64.RawURLEncoding.EncodeToString(mac.Sum(nil)) +} + +func (s *LocalStore) objectPath(scope, objectKey string) (string, error) { + if objectKey == "" || filepath.IsAbs(objectKey) { + return "", fmt.Errorf("local storage: invalid object key") + } + cleanKey := filepath.Clean(filepath.FromSlash(objectKey)) + if cleanKey == "." || cleanKey == ".." || strings.HasPrefix(cleanKey, ".."+string(filepath.Separator)) { + return "", fmt.Errorf("local storage: invalid object key") + } + root := filepath.Join(s.baseDir, scope) + target := filepath.Join(root, cleanKey) + rel, err := filepath.Rel(root, target) + if err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) { + return "", fmt.Errorf("local storage: invalid object key") + } + return target, nil +}