From cd9e358bbc2a0c330f53eafbdf5aa8c75442ca6b Mon Sep 17 00:00:00 2001 From: CoffeeKnyte <67730400+CoffeeKnyte@users.noreply.github.com> Date: Fri, 31 Jul 2026 00:46:14 +0000 Subject: [PATCH] fix(playback): fail closed on transfer-registry saturation MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `transfers.Registry` caps at 10,000 entries. Past that `Begin` returned `ErrRegistryFull`, and every call site logged at Debug and served the file anyway with its byte updates discarded. With no connection cap anywhere, one actor could pin the registry and blind download-class monitoring for everyone — which is the monitoring the abuse controls depend on. Per decision A7, saturation now fails closed and is unreachable by a single actor in the first place: - A per-user concurrent-transfer cap (`playback.max_user_concurrent_transfers`, default 24) checked before the global limit, so one actor exhausts its own budget rather than the shared registry. 24 leaves headroom for multi-connection downloaders, which open 4-8 sockets per file and take a transfer id per concurrent Range request. `MaxPerUser` is a *pointer* option because a plain int cannot distinguish "unset, use the default" from "explicitly 0, unlimited". - All five download-class call sites refuse rather than serve unmonitored: 429 `transfer_limit_exceeded` for the per-user cap, 503 `monitoring_unavailable` for a full registry, both with `Retry-After`. The handoff listed four sites; `internal/api/handlers/ebook_reader.go` was the missing fifth. - The ABS file handler now admits the transfer *before* setting Content-Disposition and the audio Content-Type, so a rejection is not mislabeled as an audio attachment. Two correctness fixes fall out of doing this properly: - `End` now decrements the per-user count only when it actually removed an entry, and drops the map entry at zero so neither counts nor warning timestamps grow unbounded. - The two download handlers registered `defer transfers.End(transfer.ID)` *before* the service called `Begin`, so a duplicate id would have removed another request's live record. `Begin` moves into the handler beside its `End`, and the service keeps its late DownloadID/MediaFileID enrichment through a new `Registry.Annotate`. Saturation is logged at Warn inside the registry, rate-limited per user and carrying the route and user id, rather than once per rejected request at each call site. A7 asked for Warn at the call sites, but an actor parked at its cap would then generate unbounded warning traffic — a log-amplification vector of its own. No information is lost: route and user id are already on the Transfer. 429/503 are new failure statuses on existing endpoints, not repurposed ones. silo-android and silo-apple need follow-up to retry with backoff on `Retry-After`; the cap and the fail-closed behavior are advertised on `GET /downloads/capability` so clients can feature-detect instead of probing. The jellycompat call site is verified by reading only: that test package does not compile on origin/main (pre-existing, unrelated). Part 3 of 3 for the Batch 4 liveness/replica work. Part of #305 --- cmd/silo/main.go | 13 +- cmd/silo/transfer_wiring.go | 4 +- cmd/silo/transfer_wiring_test.go | 5 +- internal/api/handlers/downloads.go | 64 ++++++--- internal/api/handlers/downloads_test.go | 88 +++++++++++- internal/api/handlers/ebook_reader.go | 13 +- internal/api/handlers/ebook_reader_test.go | 74 +++++++++++ internal/audiobooks/abs/file_handler.go | 53 +++++--- .../abs/file_handler_public_track_test.go | 77 +++++++++++ internal/audiobooks/abs/rss_feeds_handler.go | 4 +- .../audiobooks/abs/rss_feeds_handler_test.go | 80 +++++++++++ internal/config/admin_settings.go | 3 + internal/config/admin_settings_test.go | 10 ++ internal/config/restart_keys.go | 11 +- internal/config/restart_keys_test.go | 1 + internal/downloads/service.go | 12 +- internal/jellycompat/streams.go | 13 +- internal/transfers/registry.go | 106 ++++++++++++++- internal/transfers/registry_test.go | 125 +++++++++++++++++- 19 files changed, 697 insertions(+), 59 deletions(-) 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)