fix(playback): fail closed on transfer-registry saturation
`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
This commit is contained in:
+12
-1
@@ -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-
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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{}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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},
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user