diff --git a/cmd/silo/main.go b/cmd/silo/main.go index 1caded2c..6281508b 100644 --- a/cmd/silo/main.go +++ b/cmd/silo/main.go @@ -105,6 +105,7 @@ import ( "github.com/Silo-Server/silo-server/internal/taskmanager/triggers" "github.com/Silo-Server/silo-server/internal/telemetry" "github.com/Silo-Server/silo-server/internal/transcodenode" + "github.com/Silo-Server/silo-server/internal/transfers" "github.com/Silo-Server/silo-server/internal/usercollections" "github.com/Silo-Server/silo-server/internal/userdb" "github.com/Silo-Server/silo-server/internal/userstore" @@ -758,7 +759,17 @@ func main() { // One process-local registry is shared by every download-class serving // surface and the admin view. It is intentionally separate from live streams. - transferMonitoring := newTransferRegistryComposition() + maxUserTransfers, maxUserTransfersErr := transfers.ParseMaxPerUser(settings[transfers.MaxPerUserSetting]) + if maxUserTransfersErr != nil { + slog.Warn("invalid max user concurrent transfers setting; using default", + "component", "transfers", + "setting", transfers.MaxPerUserSetting, + "value", settings[transfers.MaxPerUserSetting], + "error", maxUserTransfersErr, + ) + maxUserTransfers, _ = transfers.ParseMaxPerUser("") + } + transferMonitoring := newTransferRegistryComposition(maxUserTransfers) // Central-side kill switch: written by the async enforcer, admin terminate, // and account revocation; edges enforce it via Redis pub/sub + poll. Memory- diff --git a/cmd/silo/transfer_wiring.go b/cmd/silo/transfer_wiring.go index e437223c..c21e2d79 100644 --- a/cmd/silo/transfer_wiring.go +++ b/cmd/silo/transfer_wiring.go @@ -13,8 +13,8 @@ type transferRegistryComposition struct { registry *transfers.Registry } -func newTransferRegistryComposition() *transferRegistryComposition { - return &transferRegistryComposition{registry: transfers.New()} +func newTransferRegistryComposition(maxPerUser int) *transferRegistryComposition { + return &transferRegistryComposition{registry: transfers.NewWithOptions(transfers.Options{MaxPerUser: &maxPerUser})} } func (c *transferRegistryComposition) wireAPI(deps *api.Dependencies) { diff --git a/cmd/silo/transfer_wiring_test.go b/cmd/silo/transfer_wiring_test.go index 0147944a..23050bba 100644 --- a/cmd/silo/transfer_wiring_test.go +++ b/cmd/silo/transfer_wiring_test.go @@ -9,7 +9,7 @@ import ( ) func TestTransferRegistryCompositionSharesOneInstance(t *testing.T) { - composition := newTransferRegistryComposition() + composition := newTransferRegistryComposition(24) var apiDeps api.Dependencies var compatDeps jellycompat.Dependencies var absDeps audiobooks.ABSHandlerDeps @@ -21,6 +21,9 @@ func TestTransferRegistryCompositionSharesOneInstance(t *testing.T) { if apiDeps.TransferRegistry == nil { t.Fatal("API/native/admin transfer registry is nil") } + if got := apiDeps.TransferRegistry.MaxPerUser(); got != 24 { + t.Fatalf("max per-user transfers = %d, want 24", got) + } if apiDeps.TransferRegistry != compatDeps.TransferRegistry { t.Fatal("jellycompat received a different transfer registry") } diff --git a/internal/api/handlers/downloads.go b/internal/api/handlers/downloads.go index 546591cb..511d198e 100644 --- a/internal/api/handlers/downloads.go +++ b/internal/api/handlers/downloads.go @@ -152,14 +152,16 @@ type batchManifestsResponse struct { // downloadCapabilityResponse is the GET /downloads/capability payload clients // use for feature detection instead of introspecting admin settings. type downloadCapabilityResponse struct { - Enabled bool `json:"enabled"` - DownloadAllowed bool `json:"download_allowed"` - QualityPresets []string `json:"quality_presets"` - TranscodeEnabled bool `json:"transcode_enabled"` - TranscodeUserAllowed bool `json:"transcode_user_allowed"` - SeasonDownload bool `json:"season_download"` - SeriesMonitoring bool `json:"series_monitoring"` - MonitoringModes []string `json:"monitoring_modes,omitempty"` + Enabled bool `json:"enabled"` + DownloadAllowed bool `json:"download_allowed"` + QualityPresets []string `json:"quality_presets"` + TranscodeEnabled bool `json:"transcode_enabled"` + TranscodeUserAllowed bool `json:"transcode_user_allowed"` + SeasonDownload bool `json:"season_download"` + SeriesMonitoring bool `json:"series_monitoring"` + MonitoringModes []string `json:"monitoring_modes,omitempty"` + MaxUserConcurrentTransfers int `json:"max_user_concurrent_transfers"` + TransferMonitoringFailClosed bool `json:"transfer_monitoring_fail_closed"` } func toDownloadResponse(d *downloads.Download) downloadResponse { @@ -216,14 +218,16 @@ func (h *DownloadHandler) HandleCapability(w http.ResponseWriter, r *http.Reques } writeJSON(w, http.StatusOK, downloadCapabilityResponse{ - Enabled: capInfo.Enabled, - DownloadAllowed: capInfo.DownloadAllowed, - QualityPresets: capInfo.QualityPresets, - TranscodeEnabled: capInfo.TranscodeEnabled, - TranscodeUserAllowed: capInfo.TranscodeUserAllowed, - SeasonDownload: capInfo.SeasonDownload, - SeriesMonitoring: capInfo.SeriesMonitoring, - MonitoringModes: capInfo.MonitoringModes, + Enabled: capInfo.Enabled, + DownloadAllowed: capInfo.DownloadAllowed, + QualityPresets: capInfo.QualityPresets, + TranscodeEnabled: capInfo.TranscodeEnabled, + TranscodeUserAllowed: capInfo.TranscodeUserAllowed, + SeasonDownload: capInfo.SeasonDownload, + SeriesMonitoring: capInfo.SeriesMonitoring, + MonitoringModes: capInfo.MonitoringModes, + MaxUserConcurrentTransfers: h.transfers.MaxPerUser(), + TransferMonitoringFailClosed: true, }) } @@ -447,6 +451,9 @@ func (h *DownloadHandler) HandleDownloadFile(w http.ResponseWriter, r *http.Requ // the write deadline with progress instead. sw := httpstream.NewRollingDeadlineWriterCtx(r.Context(), w) transfer := h.newTransfer(r, userID, profileID, deviceName, "native_download") + if !h.beginTransfer(w, r, transfer) { + return + } defer h.transfers.End(transfer.ID) metered := playback.NewSessionMeteredWriter(sw, h.transfers, transfer.ID) defer func() { _ = metered.Close() }() @@ -502,6 +509,9 @@ func (h *DownloadHandler) HandleDirectDownload(w http.ResponseWriter, r *http.Re startedAt := time.Now() r = r.WithContext(httpstream.WithCutLatch(r.Context(), &httpstream.CutLatch{})) transfer := h.newTransfer(r, userID, profileID, deviceName, "native_direct") + if !h.beginTransfer(w, r, transfer) { + return + } defer h.transfers.End(transfer.ID) metered := playback.NewSessionMeteredWriter(w, h.transfers, transfer.ID) defer func() { _ = metered.Close() }() @@ -515,6 +525,28 @@ func (h *DownloadHandler) HandleDirectDownload(w http.ResponseWriter, r *http.Re } } +func (h *DownloadHandler) beginTransfer(w http.ResponseWriter, r *http.Request, transfer transfers.Transfer) bool { + if err := h.transfers.Begin(transfer); err != nil { + slog.DebugContext(r.Context(), "download transfer rejected", + "component", "api", + "transfer_id", transfer.ID, + "error", err, + ) + switch { + case errors.Is(err, transfers.ErrUserTransferLimit): + w.Header().Set("Retry-After", "5") + writeError(w, http.StatusTooManyRequests, "transfer_limit_exceeded", "Concurrent transfer limit exceeded") + case errors.Is(err, transfers.ErrRegistryFull): + w.Header().Set("Retry-After", "5") + writeError(w, http.StatusServiceUnavailable, "monitoring_unavailable", "Transfer monitoring unavailable") + default: + writeError(w, http.StatusInternalServerError, "internal_error", "Failed to monitor transfer") + } + return false + } + return true +} + func (h *DownloadHandler) newTransfer(r *http.Request, userID int, profileID, clientName, route string) transfers.Transfer { if clientName == "" { clientName = r.UserAgent() diff --git a/internal/api/handlers/downloads_test.go b/internal/api/handlers/downloads_test.go index 7e6e6dd7..ec7ee1d1 100644 --- a/internal/api/handlers/downloads_test.go +++ b/internal/api/handlers/downloads_test.go @@ -165,7 +165,7 @@ func (f *fakeDownloadService) ServeDirect(_ context.Context, w http.ResponseWrit f.gotTransfer = transfer if f.beginTransfer && f.transferRegistry != nil { transfer.MediaFileID = fileID - _ = f.transferRegistry.Begin(transfer) + f.transferRegistry.Annotate(transfer.ID, "", fileID) } if f.directErr != nil { return f.directErr @@ -179,7 +179,7 @@ func (f *fakeDownloadService) ServeFile(_ context.Context, w http.ResponseWriter f.gotTransfer = transfer if f.beginTransfer && f.transferRegistry != nil { transfer.DownloadID = downloadID - _ = f.transferRegistry.Begin(transfer) + f.transferRegistry.Annotate(transfer.ID, downloadID, 0) } if f.serveErr != nil { return f.serveErr @@ -281,6 +281,7 @@ func TestHandleCapability(t *testing.T) { TranscodeUserAllowed: true, }} h := NewDownloadHandler(svc) + h.SetTransferRegistry(transfers.New()) rec := httptest.NewRecorder() h.HandleCapability(rec, downloadTestRequest(http.MethodGet, "/downloads/capability", nil, 7, "", "")) @@ -298,6 +299,9 @@ func TestHandleCapability(t *testing.T) { if len(resp.QualityPresets) != 1 || resp.QualityPresets[0] != downloads.QualityOriginal { t.Fatalf("quality presets = %v, want [original]", resp.QualityPresets) } + if resp.MaxUserConcurrentTransfers != 24 || !resp.TransferMonitoringFailClosed { + t.Fatalf("transfer capability = %+v", resp) + } } func TestHandleCapabilityUnauthorized(t *testing.T) { @@ -538,6 +542,86 @@ func TestHandleDirectDownloadMissingFileID(t *testing.T) { } } +func TestDownloadTransferAdmissionFailsClosed(t *testing.T) { + for _, tc := range []struct { + name string + registry func(t *testing.T) *transfers.Registry + wantStatus int + wantCode string + }{ + { + name: "per-user cap", + registry: func(t *testing.T) *transfers.Registry { + t.Helper() + limit := 1 + r := transfers.NewWithOptions(transfers.Options{MaxPerUser: &limit}) + if err := r.Begin(transfers.Transfer{ID: "existing", UserID: 7}); err != nil { + t.Fatal(err) + } + return r + }, + wantStatus: http.StatusTooManyRequests, + wantCode: "transfer_limit_exceeded", + }, + { + name: "global cap", + registry: func(t *testing.T) *transfers.Registry { + t.Helper() + unlimited := 0 + r := transfers.NewWithOptions(transfers.Options{MaxEntries: 1, MaxPerUser: &unlimited}) + if err := r.Begin(transfers.Transfer{ID: "existing", UserID: 8}); err != nil { + t.Fatal(err) + } + return r + }, + wantStatus: http.StatusServiceUnavailable, + wantCode: "monitoring_unavailable", + }, + } { + for _, route := range []string{"managed", "direct"} { + for _, method := range []string{http.MethodGet, http.MethodHead, "range"} { + t.Run(tc.name+"/"+route+"/"+method, func(t *testing.T) { + registry := tc.registry(t) + svc := &fakeDownloadService{} + h := NewDownloadHandler(svc) + h.SetTransferRegistry(registry) + requestMethod := method + if method == "range" { + requestMethod = http.MethodGet + } + var req *http.Request + if route == "managed" { + req = withChiID(downloadTestRequest(requestMethod, "/downloads/dl1/file", nil, 7, "p1", "dev1"), "dl1") + } else { + req = downloadTestRequest(requestMethod, "/direct-download?file_id=42", nil, 7, "p1", "") + } + if method == "range" { + req.Header.Set("Range", "bytes=0-1") + } + rec := httptest.NewRecorder() + if route == "managed" { + h.HandleDownloadFile(rec, req) + } else { + h.HandleDirectDownload(rec, req) + } + if rec.Code != tc.wantStatus { + t.Fatalf("status = %d, want %d; body=%s", rec.Code, tc.wantStatus, rec.Body.String()) + } + if rec.Header().Get("Retry-After") != "5" { + t.Fatalf("Retry-After = %q, want 5", rec.Header().Get("Retry-After")) + } + if !bytes.Contains(rec.Body.Bytes(), []byte(tc.wantCode)) || bytes.Contains(rec.Body.Bytes(), []byte("served")) { + t.Fatalf("body = %q, want error code and no served bytes", rec.Body.String()) + } + if svc.gotTransfer.ID != "" { + t.Fatalf("service was called with transfer %+v", svc.gotTransfer) + } + }) + } + } + } +} + func TestDownloadTransferResolutionFailureNeverRegisters(t *testing.T) { registry := transfers.New() svc := &fakeDownloadService{directErr: catalog.ErrItemNotFound} diff --git a/internal/api/handlers/ebook_reader.go b/internal/api/handlers/ebook_reader.go index fc8554d8..cfb00268 100644 --- a/internal/api/handlers/ebook_reader.go +++ b/internal/api/handlers/ebook_reader.go @@ -165,11 +165,22 @@ func (h *EbookReaderHandler) HandleReadFile(w http.ResponseWriter, r *http.Reque } if h.Transfers != nil { if err := h.Transfers.Begin(transfer); err != nil { - slog.DebugContext(r.Context(), "ebook transfer not monitored", + slog.DebugContext(r.Context(), "ebook transfer rejected", "component", "ebooks", "transfer_id", transfer.ID, "error", err, ) + switch { + case errors.Is(err, transfers.ErrUserTransferLimit): + w.Header().Set("Retry-After", "5") + writeError(w, http.StatusTooManyRequests, "transfer_limit_exceeded", "Concurrent transfer limit exceeded") + case errors.Is(err, transfers.ErrRegistryFull): + w.Header().Set("Retry-After", "5") + writeError(w, http.StatusServiceUnavailable, "monitoring_unavailable", "Transfer monitoring unavailable") + default: + writeError(w, http.StatusInternalServerError, "internal_error", "Failed to monitor transfer") + } + return } } defer h.Transfers.End(transfer.ID) diff --git a/internal/api/handlers/ebook_reader_test.go b/internal/api/handlers/ebook_reader_test.go index ade87c13..152ffc5f 100644 --- a/internal/api/handlers/ebook_reader_test.go +++ b/internal/api/handlers/ebook_reader_test.go @@ -210,6 +210,80 @@ func TestEbookReaderServesEbookInlineWithRangeSupport(t *testing.T) { } } +func TestEbookReaderTransferAdmissionFailsClosed(t *testing.T) { + for _, tc := range []struct { + name string + registry func(t *testing.T) *transfers.Registry + wantStatus int + wantCode string + }{ + { + name: "per-user cap", + registry: func(t *testing.T) *transfers.Registry { + t.Helper() + limit := 1 + r := transfers.NewWithOptions(transfers.Options{MaxPerUser: &limit}) + if err := r.Begin(transfers.Transfer{ID: "existing", UserID: 1}); err != nil { + t.Fatal(err) + } + return r + }, + wantStatus: http.StatusTooManyRequests, + wantCode: "transfer_limit_exceeded", + }, + { + name: "global cap", + registry: func(t *testing.T) *transfers.Registry { + t.Helper() + unlimited := 0 + r := transfers.NewWithOptions(transfers.Options{MaxEntries: 1, MaxPerUser: &unlimited}) + if err := r.Begin(transfers.Transfer{ID: "existing", UserID: 2}); err != nil { + t.Fatal(err) + } + return r + }, + wantStatus: http.StatusServiceUnavailable, + wantCode: "monitoring_unavailable", + }, + } { + for _, method := range []string{http.MethodGet, http.MethodHead, "range"} { + t.Run(tc.name+"/"+method, func(t *testing.T) { + registry := tc.registry(t) + handler := NewEbookReaderHandler(&MediaFileAuthorizer{ + FileResolver: stubMediaFileResolver{file: &models.MediaFile{ + ID: 42, ContentID: "ebook-1", FilePath: "must-not-be-served.epub", Container: "epub", BaseType: "ebook", + }}, + ItemAccess: stubItemAccessChecker{}, + }) + handler.Transfers = registry + requestMethod := method + if method == "range" { + requestMethod = http.MethodGet + } + req := withEbookReaderRouteParams( + newEbookReaderAuthRequest(requestMethod, "/ebooks/ebook-1/files/42/read"), + "ebook-1", + "42", + ) + if method == "range" { + req.Header.Set("Range", "bytes=0-1") + } + rec := httptest.NewRecorder() + handler.HandleReadFile(rec, req) + if rec.Code != tc.wantStatus { + t.Fatalf("status = %d, want %d; body=%s", rec.Code, tc.wantStatus, rec.Body.String()) + } + if rec.Header().Get("Retry-After") != "5" { + t.Fatalf("Retry-After = %q, want 5", rec.Header().Get("Retry-After")) + } + if !strings.Contains(rec.Body.String(), tc.wantCode) || strings.Contains(rec.Body.String(), "must-not-be-served") { + t.Fatalf("body = %q, want error code and no served bytes", rec.Body.String()) + } + }) + } + } +} + type blockingEbookWriter struct { header http.Header reached chan struct{} diff --git a/internal/audiobooks/abs/file_handler.go b/internal/audiobooks/abs/file_handler.go index 35c3ec16..aae4b713 100644 --- a/internal/audiobooks/abs/file_handler.go +++ b/internal/audiobooks/abs/file_handler.go @@ -3,6 +3,8 @@ package abs import ( "crypto/md5" "encoding/hex" + "encoding/json" + "errors" "log/slog" "net/http" "path/filepath" @@ -119,20 +121,6 @@ func (h *Handler) handleFileStream(w http.ResponseWriter, r *http.Request) { mediaFile := files[fileIdx] - // /download variant: hint the client to save rather than stream. - if strings.HasSuffix(r.URL.Path, "/download") { - filename := filepath.Base(mediaFile.FilePath) - w.Header().Set("Content-Disposition", `attachment; filename="`+filename+`"`) - } - - // Set Content-Type for audio files. ServeDirectPlay uses MimeFromExtension - // which covers video containers; we override with audio-specific MIME - // types because ABS clients pattern-match on Content-Type. - ext := strings.ToLower(filepath.Ext(mediaFile.FilePath)) - if ct := audioContentType(ext); ct != "" { - w.Header().Set("Content-Type", ct) - } - route := "abs_file_stream" if strings.HasSuffix(r.URL.Path, "/download") { route = "abs_file_download" @@ -149,10 +137,24 @@ func (h *Handler) handleFileStream(w http.ResponseWriter, r *http.Request) { } if h.deps.Transfers != nil { if err := h.deps.Transfers.Begin(transfer); err != nil { - slog.DebugContext(r.Context(), "ABS transfer not monitored", "component", "audiobooks", "transfer_id", transfer.ID, "error", err) + slog.DebugContext(r.Context(), "ABS transfer rejected", "component", "audiobooks", "transfer_id", transfer.ID, "error", err) + writeTransferBeginError(w, err) + return } } defer h.deps.Transfers.End(transfer.ID) + + // Success-only response headers must follow transfer admission so a + // rejection is not mislabeled as audio or an attachment. + if strings.HasSuffix(r.URL.Path, "/download") { + filename := filepath.Base(mediaFile.FilePath) + w.Header().Set("Content-Disposition", `attachment; filename="`+filename+`"`) + } + ext := strings.ToLower(filepath.Ext(mediaFile.FilePath)) + if ct := audioContentType(ext); ct != "" { + w.Header().Set("Content-Type", ct) + } + metered := playback.NewSessionMeteredWriter(w, h.deps.Transfers, transfer.ID) defer func() { _ = metered.Close() }() @@ -166,6 +168,27 @@ func (h *Handler) handleFileStream(w http.ResponseWriter, r *http.Request) { } } +func writeTransferBeginError(w http.ResponseWriter, err error) { + status := http.StatusInternalServerError + code := "internal_error" + message := "Failed to monitor transfer" + switch { + case errors.Is(err, transfers.ErrUserTransferLimit): + status = http.StatusTooManyRequests + code = "transfer_limit_exceeded" + message = "Concurrent transfer limit exceeded" + w.Header().Set("Retry-After", "5") + case errors.Is(err, transfers.ErrRegistryFull): + status = http.StatusServiceUnavailable + code = "monitoring_unavailable" + message = "Transfer monitoring unavailable" + w.Header().Set("Retry-After", "5") + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(status) + _ = json.NewEncoder(w).Encode(map[string]string{"code": code, "message": message}) +} + // handlePublicTrack serves audio bytes for ONE track of a playback session. // // Real ABS Android client (v2.22.0+) builds the streaming URL as diff --git a/internal/audiobooks/abs/file_handler_public_track_test.go b/internal/audiobooks/abs/file_handler_public_track_test.go index f8c051cd..0b914867 100644 --- a/internal/audiobooks/abs/file_handler_public_track_test.go +++ b/internal/audiobooks/abs/file_handler_public_track_test.go @@ -6,6 +6,7 @@ import ( "net/http/httptest" "os" "path/filepath" + "strings" "testing" "time" @@ -13,6 +14,7 @@ import ( "github.com/Silo-Server/silo-server/internal/catalog" "github.com/Silo-Server/silo-server/internal/models" + "github.com/Silo-Server/silo-server/internal/transfers" ) // fakePlaybackSessionStore is an in-memory ABSPlaybackSessionStore for the @@ -134,6 +136,81 @@ func TestHandlePublicTrack_ServesBytesForValidSession(t *testing.T) { } } +func TestFileTransferAdmissionFailsClosedBeforeSuccessHeaders(t *testing.T) { + for _, tc := range []struct { + name string + registry func(t *testing.T) *transfers.Registry + wantStatus int + wantCode string + }{ + { + name: "per-user cap", + registry: func(t *testing.T) *transfers.Registry { + t.Helper() + limit := 1 + r := transfers.NewWithOptions(transfers.Options{MaxPerUser: &limit}) + if err := r.Begin(transfers.Transfer{ID: "existing", UserID: 1}); err != nil { + t.Fatal(err) + } + return r + }, + wantStatus: http.StatusTooManyRequests, + wantCode: "transfer_limit_exceeded", + }, + { + name: "global cap", + registry: func(t *testing.T) *transfers.Registry { + t.Helper() + unlimited := 0 + r := transfers.NewWithOptions(transfers.Options{MaxEntries: 1, MaxPerUser: &unlimited}) + if err := r.Begin(transfers.Transfer{ID: "existing", UserID: 2}); err != nil { + t.Fatal(err) + } + return r + }, + wantStatus: http.StatusServiceUnavailable, + wantCode: "monitoring_unavailable", + }, + } { + for _, method := range []string{http.MethodGet, http.MethodHead, "range"} { + t.Run(tc.name+"/"+method, func(t *testing.T) { + h, _ := newPublicTrackHandler(t, "unused", "book-1", false) + h.deps.Transfers = tc.registry(t) + requestMethod := method + if method == "range" { + requestMethod = http.MethodGet + } + req := httptest.NewRequest(requestMethod, "/api/items/book-1/file/0/download", nil) + if method == "range" { + req.Header.Set("Range", "bytes=0-1") + } + rctx := chi.NewRouteContext() + rctx.URLParams.Add("libraryItemId", "book-1") + rctx.URLParams.Add("ino", "0") + ctx := context.WithValue(req.Context(), chi.RouteCtxKey, rctx) + ctx = context.WithValue(ctx, ctxKey{}, ctxAuth{UserID: "1", ProfileID: "profile-1"}) + rec := httptest.NewRecorder() + h.handleFileStream(rec, req.WithContext(ctx)) + if rec.Code != tc.wantStatus { + t.Fatalf("status = %d, want %d; body=%s", rec.Code, tc.wantStatus, rec.Body.String()) + } + if rec.Header().Get("Retry-After") != "5" { + t.Fatalf("Retry-After = %q, want 5", rec.Header().Get("Retry-After")) + } + if rec.Header().Get("Content-Disposition") != "" { + t.Fatalf("rejection has Content-Disposition %q", rec.Header().Get("Content-Disposition")) + } + if rec.Header().Get("Content-Type") != "application/json" { + t.Fatalf("Content-Type = %q, want application/json", rec.Header().Get("Content-Type")) + } + if !strings.Contains(rec.Body.String(), tc.wantCode) || strings.Contains(rec.Body.String(), "audio-bytes") { + t.Fatalf("body = %q, want error code and no audio bytes", rec.Body.String()) + } + }) + } + } +} + // TestHandlePublicTrack_HeadProbe covers the iOS/Android HEAD pre-flight // some players issue before the GET. http.ServeContent returns headers // without a body for HEAD; the handler must not 404. diff --git a/internal/audiobooks/abs/rss_feeds_handler.go b/internal/audiobooks/abs/rss_feeds_handler.go index 5d98ec46..49b4c8d6 100644 --- a/internal/audiobooks/abs/rss_feeds_handler.go +++ b/internal/audiobooks/abs/rss_feeds_handler.go @@ -260,7 +260,9 @@ func (h *Handler) handlePublicFeedFile(w http.ResponseWriter, r *http.Request) { } if h.deps.Transfers != nil { if err := h.deps.Transfers.Begin(transfer); err != nil { - slog.DebugContext(r.Context(), "ABS feed transfer not monitored", "component", "audiobooks", "transfer_id", transfer.ID, "error", err) + slog.DebugContext(r.Context(), "ABS feed transfer rejected", "component", "audiobooks", "transfer_id", transfer.ID, "error", err) + writeTransferBeginError(w, err) + return } } defer h.deps.Transfers.End(transfer.ID) diff --git a/internal/audiobooks/abs/rss_feeds_handler_test.go b/internal/audiobooks/abs/rss_feeds_handler_test.go index e1d425a0..310eea38 100644 --- a/internal/audiobooks/abs/rss_feeds_handler_test.go +++ b/internal/audiobooks/abs/rss_feeds_handler_test.go @@ -5,13 +5,17 @@ import ( "encoding/json" "errors" "net/http" + "net/http/httptest" "sort" "strings" "sync" "testing" "time" + "github.com/go-chi/chi/v5" + "github.com/Silo-Server/silo-server/internal/models" + "github.com/Silo-Server/silo-server/internal/transfers" ) type memRSSFeedStore struct { @@ -219,3 +223,79 @@ func TestPublicFeed_HappyPath_XML(t *testing.T) { } } } + +func TestFeedFileTransferAdmissionFailsClosed(t *testing.T) { + for _, tc := range []struct { + name string + registry func(t *testing.T) *transfers.Registry + wantStatus int + wantCode string + }{ + { + name: "per-user cap", + registry: func(t *testing.T) *transfers.Registry { + t.Helper() + limit := 1 + r := transfers.NewWithOptions(transfers.Options{MaxPerUser: &limit}) + if err := r.Begin(transfers.Transfer{ID: "existing", UserID: 1}); err != nil { + t.Fatal(err) + } + return r + }, + wantStatus: http.StatusTooManyRequests, + wantCode: "transfer_limit_exceeded", + }, + { + name: "global cap", + registry: func(t *testing.T) *transfers.Registry { + t.Helper() + unlimited := 0 + r := transfers.NewWithOptions(transfers.Options{MaxEntries: 1, MaxPerUser: &unlimited}) + if err := r.Begin(transfers.Transfer{ID: "existing", UserID: 2}); err != nil { + t.Fatal(err) + } + return r + }, + wantStatus: http.StatusServiceUnavailable, + wantCode: "monitoring_unavailable", + }, + } { + for _, method := range []string{http.MethodGet, http.MethodHead, "range"} { + t.Run(tc.name+"/"+method, func(t *testing.T) { + store := newMemRSSFeedStore() + store.rows["feed-id"] = RSSFeed{ + ID: "feed-id", UserID: "1", ProfileID: "profile-1", LibraryItemID: "book-1", Slug: "feed", CreatedAt: time.Now(), + } + h := New(Dependencies{ + RSSFeedStore: store, + MediaStore: &revocationFeedFileMediaStore{file: &models.MediaFile{ + ID: 9, ContentID: "book-1", FilePath: "must-not-be-served.mp3", + }}, + Transfers: tc.registry(t), + }) + requestMethod := method + if method == "range" { + requestMethod = http.MethodGet + } + req := httptest.NewRequest(requestMethod, "/feed/feed/file/9", nil) + if method == "range" { + req.Header.Set("Range", "bytes=0-1") + } + rctx := chi.NewRouteContext() + rctx.URLParams.Add("slug", "feed") + rctx.URLParams.Add("ino", "9") + rec := httptest.NewRecorder() + h.handlePublicFeedFile(rec, req.WithContext(context.WithValue(req.Context(), chi.RouteCtxKey, rctx))) + if rec.Code != tc.wantStatus { + t.Fatalf("status = %d, want %d; body=%s", rec.Code, tc.wantStatus, rec.Body.String()) + } + if rec.Header().Get("Retry-After") != "5" { + t.Fatalf("Retry-After = %q, want 5", rec.Header().Get("Retry-After")) + } + if !strings.Contains(rec.Body.String(), tc.wantCode) || strings.Contains(rec.Body.String(), "must-not-be-served") { + t.Fatalf("body = %q, want error code and no file bytes", rec.Body.String()) + } + }) + } + } +} diff --git a/internal/config/admin_settings.go b/internal/config/admin_settings.go index 3e3e5c45..86f4e884 100644 --- a/internal/config/admin_settings.go +++ b/internal/config/admin_settings.go @@ -60,6 +60,7 @@ var adminSettingDefaults = map[string]string{ "playback.chapter_thumbnail_node_capacity": "1", "playback.chapter_thumbnail_hdr_policy": "best_effort", "playback.over_cap_revocation_ttl": streamtoken.MaxTTL.String(), + "playback.max_user_concurrent_transfers": "24", "playback.watched_threshold": "90", "playback.min_resume_threshold": "5", "allow_4k_transcode": "false", @@ -363,6 +364,8 @@ func NormalizeAdminSetting(key, raw string) (string, error) { return normalizeAdminDuration(key, value) case "playback.over_cap_revocation_ttl": return normalizeAdminDurationRange(key, value, 5*time.Minute, streamtoken.MaxTTL) + case "playback.max_user_concurrent_transfers": + return normalizeAdminInt(key, value, 0, 10_000) case "server.log_level": return normalizeAdminEnum(key, value, "debug", "info", "warn", "error") diff --git a/internal/config/admin_settings_test.go b/internal/config/admin_settings_test.go index cac413d3..89bc918d 100644 --- a/internal/config/admin_settings_test.go +++ b/internal/config/admin_settings_test.go @@ -212,6 +212,8 @@ func TestNormalizeAdminSettingRejectsInvalidValues(t *testing.T) { {key: "auth.access_token_expiry", value: "forever"}, {key: "playback.over_cap_revocation_ttl", value: "4m59s"}, {key: "playback.over_cap_revocation_ttl", value: "24h1s"}, + {key: "playback.max_user_concurrent_transfers", value: "-1"}, + {key: "playback.max_user_concurrent_transfers", value: "10001"}, {key: "recommendations.embeddings_cron", value: "not a cron"}, {key: "notifications.server_channels.batch_seconds", value: "119"}, {key: "catalog.search.meilisearch.semantic_ratio", value: "1.2"}, @@ -237,6 +239,14 @@ func TestOverCapRevocationTTLDefaultAndValidation(t *testing.T) { if got, err := NormalizeAdminSetting("playback.over_cap_revocation_ttl", "6h"); err != nil || got != "6h" { t.Fatalf("valid over-cap revocation TTL = %q, %v", got, err) } + if got := effective["playback.max_user_concurrent_transfers"]; got != "24" { + t.Fatalf("default max user concurrent transfers = %q, want 24", got) + } + for _, value := range []string{"0", "48", "10000"} { + if got, err := NormalizeAdminSetting("playback.max_user_concurrent_transfers", value); err != nil || got != value { + t.Fatalf("max user concurrent transfers %q = %q, %v", value, got, err) + } + } } func TestNormalizeAdminSettingAcceptsApprovedThemeCatalogURL(t *testing.T) { diff --git a/internal/config/restart_keys.go b/internal/config/restart_keys.go index 2130f6a3..5e4cffa2 100644 --- a/internal/config/restart_keys.go +++ b/internal/config/restart_keys.go @@ -36,11 +36,12 @@ var restartRequiredKeys = map[string]bool{ // thumbnails, audiobook enricher) — keep restart-required until those // convert. transcode_dir is fully live (only the playback handler reads // it). The chapter-thumbnail worker pool is sized at construction. - "playback.ffmpeg_path": true, - "playback.hw_accel": true, - "playback.hw_device": true, - "playback.chapter_thumbnail_workers": true, - "playback.over_cap_revocation_ttl": true, + "playback.ffmpeg_path": true, + "playback.hw_accel": true, + "playback.hw_device": true, + "playback.chapter_thumbnail_workers": true, + "playback.over_cap_revocation_ttl": true, + "playback.max_user_concurrent_transfers": true, // Scanner / matcher toggles captured at construction. Worker counts, // batch size, metadata.cache_images, and mdblist.api_key hot-reload via diff --git a/internal/config/restart_keys_test.go b/internal/config/restart_keys_test.go index 3f84e438..22e2daec 100644 --- a/internal/config/restart_keys_test.go +++ b/internal/config/restart_keys_test.go @@ -12,6 +12,7 @@ func TestRestartRequired(t *testing.T) { {"auth.jwt_secret", true}, {"ratelimit.backend", true}, {"playback.over_cap_revocation_ttl", true}, + {"playback.max_user_concurrent_transfers", true}, // Prefix-covered namespaces. {"database.max_connections", true}, {"userdb.backend", true}, diff --git a/internal/downloads/service.go b/internal/downloads/service.go index 388bf262..9d6f7921 100644 --- a/internal/downloads/service.go +++ b/internal/downloads/service.go @@ -874,7 +874,7 @@ func (s *Service) ServeDirect(ctx context.Context, w http.ResponseWriter, r *htt return catalog.ErrItemNotFound } transfer.MediaFileID = file.ID - s.beginTransfer(ctx, transfer) + s.annotateTransfer(transfer) return s.serveLocalFile(ctx, w, r, file.FilePath, userID) } @@ -1081,7 +1081,7 @@ func (s *Service) serveDownloadBytes(ctx context.Context, w http.ResponseWriter, } transfer.DownloadID = dl.ID transfer.MediaFileID = dl.MediaFileID - s.beginTransfer(ctx, transfer) + s.annotateTransfer(transfer) return s.serveLocalFile(ctx, w, r, artifact.OutputPath, userID) } if file.MissingSince != nil { @@ -1092,17 +1092,15 @@ func (s *Service) serveDownloadBytes(ctx context.Context, w http.ResponseWriter, } transfer.DownloadID = dl.ID transfer.MediaFileID = dl.MediaFileID - s.beginTransfer(ctx, transfer) + s.annotateTransfer(transfer) return s.serveLocalFile(ctx, w, r, file.FilePath, userID) } -func (s *Service) beginTransfer(ctx context.Context, transfer transfers.Transfer) { +func (s *Service) annotateTransfer(transfer transfers.Transfer) { if s.transfers == nil { return } - if err := s.transfers.Begin(transfer); err != nil { - slog.DebugContext(ctx, "download transfer not monitored", "component", "downloads", "transfer_id", transfer.ID, "error", err) - } + s.transfers.Annotate(transfer.ID, transfer.DownloadID, transfer.MediaFileID) } func (s *Service) serveLocalFile(ctx context.Context, w http.ResponseWriter, r *http.Request, path string, userID int) error { diff --git a/internal/jellycompat/streams.go b/internal/jellycompat/streams.go index 779aa673..6acb75f6 100644 --- a/internal/jellycompat/streams.go +++ b/internal/jellycompat/streams.go @@ -268,7 +268,18 @@ func (h *PlaybackHandler) HandleDownload(w http.ResponseWriter, r *http.Request) } if h.Transfers != nil { if err := h.Transfers.Begin(transfer); err != nil { - slog.DebugContext(r.Context(), "compat transfer not monitored", "component", "jellycompat", "transfer_id", transfer.ID, "error", err) + slog.DebugContext(r.Context(), "compat transfer rejected", "component", "jellycompat", "transfer_id", transfer.ID, "error", err) + switch { + case errors.Is(err, transfers.ErrUserTransferLimit): + w.Header().Set("Retry-After", "5") + writeError(w, http.StatusTooManyRequests, "transfer_limit_exceeded", "Concurrent transfer limit exceeded") + case errors.Is(err, transfers.ErrRegistryFull): + w.Header().Set("Retry-After", "5") + writeError(w, http.StatusServiceUnavailable, "monitoring_unavailable", "Transfer monitoring unavailable") + default: + writeError(w, http.StatusInternalServerError, "internal_error", "Failed to monitor transfer") + } + return } } defer h.Transfers.End(transfer.ID) diff --git a/internal/transfers/registry.go b/internal/transfers/registry.go index 8710e778..eca23d91 100644 --- a/internal/transfers/registry.go +++ b/internal/transfers/registry.go @@ -3,9 +3,11 @@ package transfers import ( "errors" + "fmt" "log/slog" "math" "sort" + "strconv" "strings" "sync" "time" @@ -14,6 +16,7 @@ import ( const ( defaultMaxEntries = 10_000 + defaultMaxPerUser = 24 fullWarningInterval = time.Minute maxDownloadIDLength = 256 maxProfileIDLength = 256 @@ -23,11 +26,14 @@ const ( ) var ( - ErrRegistryFull = errors.New("transfer registry is full") - ErrInvalidID = errors.New("transfer id is required") - ErrDuplicateID = errors.New("transfer id is already active") + ErrRegistryFull = errors.New("transfer registry is full") + ErrUserTransferLimit = errors.New("user transfer limit reached") + ErrInvalidID = errors.New("transfer id is required") + ErrDuplicateID = errors.New("transfer id is already active") ) +const MaxPerUserSetting = "playback.max_user_concurrent_transfers" + // Transfer describes one active HTTP file pour. DownloadID is correlation // metadata only: ID uniquely identifies this request, including concurrent // Range requests for the same download row. @@ -48,15 +54,25 @@ type Transfer struct { // Options customizes a Registry. Zero values use production defaults. type Options struct { MaxEntries int + MaxPerUser *int Now func() time.Time + Logger *slog.Logger +} + +type userTransferState struct { + count int + lastWarning time.Time } // Registry is a bounded, process-local collection of active transfers. type Registry struct { mu sync.RWMutex items map[string]Transfer + perUser map[int]userTransferState maxEntries int + maxPerUser int now func() time.Time + logger *slog.Logger lastFullWarning time.Time } @@ -73,13 +89,46 @@ func NewWithOptions(opts Options) *Registry { if opts.Now == nil { opts.Now = time.Now } + maxPerUser := defaultMaxPerUser + if opts.MaxPerUser != nil { + maxPerUser = max(0, *opts.MaxPerUser) + } + if opts.Logger == nil { + opts.Logger = slog.Default() + } return &Registry{ items: make(map[string]Transfer), + perUser: make(map[int]userTransferState), maxEntries: opts.MaxEntries, + maxPerUser: maxPerUser, now: opts.Now, + logger: opts.Logger, } } +// ParseMaxPerUser parses the process-start setting. Empty input selects the +// production default; zero explicitly disables the per-user cap. +func ParseMaxPerUser(raw string) (int, error) { + raw = strings.TrimSpace(raw) + if raw == "" { + return defaultMaxPerUser, nil + } + value, err := strconv.Atoi(raw) + if err != nil || value < 0 || value > defaultMaxEntries { + return 0, fmt.Errorf("%s must be an integer between 0 and %d", MaxPerUserSetting, defaultMaxEntries) + } + return value, nil +} + +// MaxPerUser returns the configured per-user concurrent-transfer cap. Zero +// means unlimited. +func (r *Registry) MaxPerUser() int { + if r == nil { + return 0 + } + return r.maxPerUser +} + // Begin records an active transfer. Strings derived from request metadata are // normalized and clamped so the entry limit also provides a useful memory bound. func (r *Registry) Begin(t Transfer) error { @@ -107,18 +156,56 @@ func (r *Registry) Begin(t Transfer) error { if _, exists := r.items[t.ID]; exists { return ErrDuplicateID } + userState := r.perUser[t.UserID] + if r.maxPerUser > 0 && userState.count >= r.maxPerUser { + now := r.now() + if userState.lastWarning.IsZero() || now.Sub(userState.lastWarning) >= fullWarningInterval { + userState.lastWarning = now + r.perUser[t.UserID] = userState + r.logger.Warn("user transfer limit reached", + "component", "transfers", + "route", t.Route, + "user_id", t.UserID, + "max_per_user", r.maxPerUser, + ) + } + return ErrUserTransferLimit + } if len(r.items) >= r.maxEntries { now := r.now() if r.lastFullWarning.IsZero() || now.Sub(r.lastFullWarning) >= fullWarningInterval { r.lastFullWarning = now - slog.Warn("transfer registry full; serving without monitoring", "component", "transfers", "max_entries", r.maxEntries) + r.logger.Warn("transfer registry full", + "component", "transfers", + "route", t.Route, + "user_id", t.UserID, + "max_entries", r.maxEntries, + ) } return ErrRegistryFull } r.items[t.ID] = t + userState.count++ + r.perUser[t.UserID] = userState return nil } +// Annotate adds correlation metadata after a handler-owned Begin. Unknown IDs +// are ignored so late resolution cannot create an unbounded registry entry. +func (r *Registry) Annotate(id, downloadID string, mediaFileID int) { + if r == nil { + return + } + r.mu.Lock() + t, ok := r.items[id] + if ok { + t.DownloadID = clamp(downloadID, maxDownloadIDLength) + t.MediaFileID = mediaFileID + r.items[id] = t + } + r.mu.Unlock() +} + // AddServedBytes implements playback.ServedBytesRecorder. Updates for unknown // transfers are deliberately discarded, and non-positive values are ignored. // The metered writer reports in coarse chunks (currently 1 MiB and on Close), @@ -149,7 +236,16 @@ func (r *Registry) End(id string) { return } r.mu.Lock() - delete(r.items, id) + if t, ok := r.items[id]; ok { + delete(r.items, id) + state := r.perUser[t.UserID] + state.count-- + if state.count <= 0 { + delete(r.perUser, t.UserID) + } else { + r.perUser[t.UserID] = state + } + } r.mu.Unlock() } diff --git a/internal/transfers/registry_test.go b/internal/transfers/registry_test.go index 8182d6fa..cf317b53 100644 --- a/internal/transfers/registry_test.go +++ b/internal/transfers/registry_test.go @@ -1,12 +1,15 @@ package transfers import ( + "context" "errors" "io" + "log/slog" "net/http" "net/http/httptest" "os" "path/filepath" + "strconv" "strings" "sync" "sync/atomic" @@ -82,7 +85,8 @@ func TestRegistryConcurrentPoursForSameDownloadDoNotCollide(t *testing.T) { } func TestRegistryConcurrentLifecycle(t *testing.T) { - r := NewWithOptions(Options{MaxEntries: 256}) + unlimited := 0 + r := NewWithOptions(Options{MaxEntries: 256, MaxPerUser: &unlimited}) var wg sync.WaitGroup for i := 0; i < 200; i++ { wg.Add(1) @@ -108,7 +112,8 @@ func TestRegistryConcurrentLifecycle(t *testing.T) { func TestRegistryCapHoldsUnderRacingBegin(t *testing.T) { const cap = 7 - r := NewWithOptions(Options{MaxEntries: cap}) + unlimited := 0 + r := NewWithOptions(Options{MaxEntries: cap, MaxPerUser: &unlimited}) var wg sync.WaitGroup var accepted atomic.Int64 for i := 0; i < 100; i++ { @@ -132,6 +137,122 @@ func TestRegistryCapHoldsUnderRacingBegin(t *testing.T) { } } +func TestRegistryPerUserLimitAndAnnotation(t *testing.T) { + now := time.Date(2026, 7, 31, 0, 0, 0, 0, time.UTC) + maxPerUser := 2 + logs := &countingLogHandler{} + r := NewWithOptions(Options{ + MaxPerUser: &maxPerUser, + Now: func() time.Time { return now }, + Logger: slog.New(logs), + }) + for _, id := range []string{"one", "two"} { + if err := r.Begin(Transfer{ID: id, UserID: 7, Route: "native_download"}); err != nil { + t.Fatalf("Begin(%q): %v", id, err) + } + } + if err := r.Begin(Transfer{ID: "over", UserID: 7, Route: "native_download"}); !errors.Is(err, ErrUserTransferLimit) { + t.Fatalf("Begin at per-user cap = %v, want ErrUserTransferLimit", err) + } + if err := r.Begin(Transfer{ID: "other", UserID: 8}); err != nil { + t.Fatalf("other user Begin: %v", err) + } + + r.End("never-began") + if err := r.Begin(Transfer{ID: "still-over", UserID: 7}); !errors.Is(err, ErrUserTransferLimit) { + t.Fatalf("Begin after unknown End = %v, want ErrUserTransferLimit", err) + } + if got := logs.count.Load(); got != 1 { + t.Fatalf("warnings before interval = %d, want 1", got) + } + now = now.Add(fullWarningInterval) + if err := r.Begin(Transfer{ID: "warn-again", UserID: 7}); !errors.Is(err, ErrUserTransferLimit) { + t.Fatalf("Begin after warning interval = %v, want ErrUserTransferLimit", err) + } + if got := logs.count.Load(); got != 2 { + t.Fatalf("warnings after interval = %d, want 2", got) + } + + r.End("one") + if err := r.Begin(Transfer{ID: "replacement", UserID: 7}); err != nil { + t.Fatalf("Begin after End: %v", err) + } + r.Annotate("replacement", " download ", 42) + r.Annotate("unknown", "ignored", 99) + var replacement Transfer + for _, transfer := range r.Snapshot() { + if transfer.ID == "replacement" { + replacement = transfer + } + } + if replacement.DownloadID != "download" || replacement.MediaFileID != 42 { + t.Fatalf("annotated transfer = %+v", replacement) + } + if len(r.Snapshot()) != 3 { + t.Fatalf("Annotate unknown changed registry: %+v", r.Snapshot()) + } + + r.End("two") + r.End("replacement") + r.End("other") + if len(r.perUser) != 0 { + t.Fatalf("perUser retained zero-count state: %+v", r.perUser) + } + + unlimited := 0 + unlimitedRegistry := NewWithOptions(Options{MaxEntries: 30, MaxPerUser: &unlimited}) + for i := 0; i < 30; i++ { + if err := unlimitedRegistry.Begin(Transfer{ID: strconv.Itoa(i), UserID: 1}); err != nil { + t.Fatalf("explicit unlimited Begin(%d): %v", i, err) + } + } + + defaultRegistry := New() + if defaultRegistry.MaxPerUser() != defaultMaxPerUser { + t.Fatalf("unset MaxPerUser = %d, want %d", defaultRegistry.MaxPerUser(), defaultMaxPerUser) + } + for i := 0; i < defaultMaxPerUser; i++ { + if err := defaultRegistry.Begin(Transfer{ID: strconv.Itoa(i), UserID: 1}); err != nil { + t.Fatalf("default Begin(%d): %v", i, err) + } + } + if err := defaultRegistry.Begin(Transfer{ID: "default-over", UserID: 1}); !errors.Is(err, ErrUserTransferLimit) { + t.Fatalf("Begin at default cap = %v, want ErrUserTransferLimit", err) + } +} + +func TestParseMaxPerUser(t *testing.T) { + for _, tc := range []struct { + raw string + want int + wantErr bool + }{ + {raw: "", want: defaultMaxPerUser}, + {raw: "0", want: 0}, + {raw: "48", want: 48}, + {raw: "-1", wantErr: true}, + {raw: "10001", wantErr: true}, + {raw: "many", wantErr: true}, + } { + got, err := ParseMaxPerUser(tc.raw) + if (err != nil) != tc.wantErr || got != tc.want { + t.Errorf("ParseMaxPerUser(%q) = %d, %v; want %d, error=%v", tc.raw, got, err, tc.want, tc.wantErr) + } + } +} + +type countingLogHandler struct { + count atomic.Int64 +} + +func (h *countingLogHandler) Enabled(context.Context, slog.Level) bool { return true } +func (h *countingLogHandler) Handle(context.Context, slog.Record) error { + h.count.Add(1) + return nil +} +func (h *countingLogHandler) WithAttrs([]slog.Attr) slog.Handler { return h } +func (h *countingLogHandler) WithGroup(string) slog.Handler { return h } + func TestRegistryClampsAndNormalizesRequestStrings(t *testing.T) { r := New() long := strings.Repeat("界", 400)