Files
at-container-registry/pkg/hold/multipart.go
T

337 lines
10 KiB
Go

package hold
import (
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"log"
"sync"
"time"
"github.com/google/uuid"
)
// MultipartMode indicates how multipart uploads are handled
type MultipartMode int
const (
// S3Native uses S3's native multipart API with presigned URLs
S3Native MultipartMode = iota
// Buffered buffers parts in memory and assembles them in the hold service
Buffered
)
// CompletedPart represents an uploaded part with its ETag
type CompletedPart struct {
PartNumber int `json:"part_number"`
ETag string `json:"etag"`
}
// MultipartSession tracks an in-progress multipart upload
type MultipartSession struct {
UploadID string // Unique upload ID
Digest string // Target digest path
Mode MultipartMode // Upload mode (S3Native or Buffered)
S3UploadID string // S3 upload ID (for S3Native mode)
Parts map[int]*MultipartPart // Buffered parts (for Buffered mode)
CreatedAt time.Time // When upload started
LastActivity time.Time // Last part upload
mu sync.RWMutex // Protects Parts map
}
// MultipartPart represents a single part in a multipart upload
type MultipartPart struct {
PartNumber int // Part number (1-indexed)
Data []byte // Part data (for Buffered mode)
ETag string // ETag from S3 or computed hash
Size int64 // Part size in bytes
UploadedAt time.Time // When part was uploaded
}
// MultipartManager manages multipart upload sessions
type MultipartManager struct {
sessions map[string]*MultipartSession // uploadID -> session
mu sync.RWMutex // Protects sessions map
}
// NewMultipartManager creates a new multipart manager
func NewMultipartManager() *MultipartManager {
mgr := &MultipartManager{
sessions: make(map[string]*MultipartSession),
}
// Start cleanup goroutine for abandoned uploads
go mgr.cleanupLoop()
return mgr
}
// cleanupLoop periodically cleans up expired sessions
func (m *MultipartManager) cleanupLoop() {
ticker := time.NewTicker(15 * time.Minute)
defer ticker.Stop()
for range ticker.C {
m.cleanupExpiredSessions()
}
}
// cleanupExpiredSessions removes sessions inactive for >24 hours
func (m *MultipartManager) cleanupExpiredSessions() {
m.mu.Lock()
defer m.mu.Unlock()
now := time.Now()
for uploadID, session := range m.sessions {
if now.Sub(session.LastActivity) > 24*time.Hour {
log.Printf("Cleaning up expired multipart session: uploadID=%s, age=%v", uploadID, now.Sub(session.CreatedAt))
delete(m.sessions, uploadID)
}
}
}
// CreateSession creates a new multipart upload session
func (m *MultipartManager) CreateSession(digest string, mode MultipartMode, s3UploadID string) *MultipartSession {
uploadID := uuid.New().String()
session := &MultipartSession{
UploadID: uploadID,
Digest: digest,
Mode: mode,
S3UploadID: s3UploadID,
Parts: make(map[int]*MultipartPart),
CreatedAt: time.Now(),
LastActivity: time.Now(),
}
m.mu.Lock()
m.sessions[uploadID] = session
m.mu.Unlock()
log.Printf("Created multipart session: uploadID=%s, digest=%s, mode=%v", uploadID, digest, mode)
return session
}
// GetSession retrieves a multipart session by upload ID
func (m *MultipartManager) GetSession(uploadID string) (*MultipartSession, error) {
m.mu.RLock()
defer m.mu.RUnlock()
session, ok := m.sessions[uploadID]
if !ok {
return nil, fmt.Errorf("multipart session not found: %s", uploadID)
}
return session, nil
}
// DeleteSession removes a multipart session
func (m *MultipartManager) DeleteSession(uploadID string) {
m.mu.Lock()
defer m.mu.Unlock()
delete(m.sessions, uploadID)
log.Printf("Deleted multipart session: uploadID=%s", uploadID)
}
// StorePart stores a part in the session (for Buffered mode)
func (s *MultipartSession) StorePart(partNumber int, data []byte) string {
s.mu.Lock()
defer s.mu.Unlock()
// Compute ETag as SHA256 hash of part data
hash := sha256.Sum256(data)
etag := hex.EncodeToString(hash[:])
part := &MultipartPart{
PartNumber: partNumber,
Data: data,
ETag: etag,
Size: int64(len(data)),
UploadedAt: time.Now(),
}
s.Parts[partNumber] = part
s.LastActivity = time.Now()
log.Printf("Stored part: uploadID=%s, part=%d, size=%d bytes, etag=%s", s.UploadID, partNumber, len(data), etag)
return etag
}
// RecordS3Part records a part uploaded to S3 (for S3Native mode)
func (s *MultipartSession) RecordS3Part(partNumber int, etag string, size int64) {
s.mu.Lock()
defer s.mu.Unlock()
part := &MultipartPart{
PartNumber: partNumber,
ETag: etag,
Size: size,
UploadedAt: time.Now(),
}
s.Parts[partNumber] = part
s.LastActivity = time.Now()
log.Printf("Recorded S3 part: uploadID=%s, part=%d, size=%d bytes, etag=%s", s.UploadID, partNumber, size, etag)
}
// AssembleBufferedParts assembles all buffered parts into a single blob
// Returns the complete data and total size
func (s *MultipartSession) AssembleBufferedParts() ([]byte, int64, error) {
s.mu.RLock()
defer s.mu.RUnlock()
if s.Mode != Buffered {
return nil, 0, fmt.Errorf("session is not in buffered mode")
}
// Calculate total size
var totalSize int64
maxPart := 0
for partNum, part := range s.Parts {
totalSize += part.Size
if partNum > maxPart {
maxPart = partNum
}
}
// Check for missing parts
for i := 1; i <= maxPart; i++ {
if _, ok := s.Parts[i]; !ok {
return nil, 0, fmt.Errorf("missing part %d", i)
}
}
// Assemble parts in order
assembled := make([]byte, 0, totalSize)
for i := 1; i <= maxPart; i++ {
part := s.Parts[i]
assembled = append(assembled, part.Data...)
}
log.Printf("Assembled buffered parts: uploadID=%s, parts=%d, totalSize=%d bytes", s.UploadID, maxPart, totalSize)
return assembled, totalSize, nil
}
// GetCompletedParts returns the list of completed parts for S3 multipart completion
func (s *MultipartSession) GetCompletedParts() []CompletedPart {
s.mu.RLock()
defer s.mu.RUnlock()
parts := make([]CompletedPart, 0, len(s.Parts))
for _, part := range s.Parts {
parts = append(parts, CompletedPart{
PartNumber: part.PartNumber,
ETag: part.ETag,
})
}
return parts
}
// StartMultipartUploadWithManager initiates a multipart upload using the manager
// Returns uploadID and mode
func (s *HoldService) StartMultipartUploadWithManager(ctx context.Context, digest string, manager *MultipartManager) (string, MultipartMode, error) {
// Check if presigned URLs are disabled for testing
if s.config.Server.DisablePresignedURLs {
log.Printf("Presigned URLs disabled (DISABLE_PRESIGNED_URLS=true), using buffered mode")
session := manager.CreateSession(digest, Buffered, "")
log.Printf("Started buffered multipart: uploadID=%s", session.UploadID)
return session.UploadID, Buffered, nil
}
// Try S3 native multipart first
if s.s3Client != nil {
s3UploadID, err := s.startMultipartUpload(ctx, digest)
if err == nil {
// S3 native multipart succeeded
session := manager.CreateSession(digest, S3Native, s3UploadID)
log.Printf("Started S3 native multipart: uploadID=%s, s3UploadID=%s", session.UploadID, s3UploadID)
return session.UploadID, S3Native, nil
}
log.Printf("S3 native multipart failed, falling back to buffered mode: %v", err)
}
// Fallback to buffered mode
session := manager.CreateSession(digest, Buffered, "")
log.Printf("Started buffered multipart: uploadID=%s", session.UploadID)
return session.UploadID, Buffered, nil
}
// GetPartUploadURL generates a presigned URL for uploading a part
// Only used for S3Native mode - Buffered mode is handled by blobstore adapter
func (s *HoldService) GetPartUploadURL(ctx context.Context, session *MultipartSession, partNumber int, did string) (string, error) {
if session.Mode != S3Native {
return "", fmt.Errorf("GetPartUploadURL only supports S3Native mode")
}
// Generate S3 presigned URL for this part
url, err := s.getPartPresignedURL(ctx, session.Digest, session.S3UploadID, partNumber)
if err != nil {
return "", fmt.Errorf("failed to generate S3 part URL: %w", err)
}
return url, nil
}
// CompleteMultipartUploadWithManager completes a multipart upload
func (s *HoldService) CompleteMultipartUploadWithManager(ctx context.Context, session *MultipartSession, manager *MultipartManager) error {
defer manager.DeleteSession(session.UploadID)
if session.Mode == S3Native {
// Complete S3 multipart upload
parts := session.GetCompletedParts()
if err := s.completeMultipartUpload(ctx, session.Digest, session.S3UploadID, parts); err != nil {
return fmt.Errorf("failed to complete S3 multipart: %w", err)
}
log.Printf("Completed S3 native multipart: uploadID=%s, parts=%d", session.UploadID, len(parts))
return nil
}
// Buffered mode: assemble parts and write via driver
data, size, err := session.AssembleBufferedParts()
if err != nil {
return fmt.Errorf("failed to assemble parts: %w", err)
}
// Write assembled blob to storage
path := blobPath(session.Digest)
writer, err := s.driver.Writer(ctx, path, false)
if err != nil {
return fmt.Errorf("failed to create writer: %w", err)
}
written, err := writer.Write(data)
if err != nil {
writer.Cancel(ctx)
return fmt.Errorf("failed to write blob: %w", err)
}
if err := writer.Commit(ctx); err != nil {
return fmt.Errorf("failed to commit blob: %w", err)
}
log.Printf("Completed buffered multipart: uploadID=%s, size=%d bytes, written=%d", session.UploadID, size, written)
return nil
}
// AbortMultipartUploadWithManager aborts a multipart upload
func (s *HoldService) AbortMultipartUploadWithManager(ctx context.Context, session *MultipartSession, manager *MultipartManager) error {
defer manager.DeleteSession(session.UploadID)
if session.Mode == S3Native {
// Abort S3 multipart upload
if err := s.abortMultipartUpload(ctx, session.Digest, session.S3UploadID); err != nil {
return fmt.Errorf("failed to abort S3 multipart: %w", err)
}
log.Printf("Aborted S3 native multipart: uploadID=%s", session.UploadID)
return nil
}
// Buffered mode: just delete the session (parts are in memory)
log.Printf("Aborted buffered multipart: uploadID=%s", session.UploadID)
return nil
}