package handler import ( "context" "net/http" "net/http/httptest" "net/url" "testing" "github.com/go-chi/chi/v5" ) // A flash must be shown exactly once: reading it also expires the cookie. func TestConsumeFlashClearsCookie(t *testing.T) { w := httptest.NewRecorder() setFlash(w, "success", "Preset created.") cookie := w.Result().Cookies()[0] r := httptest.NewRequest("GET", "/settings", nil) r.AddCookie(cookie) next := httptest.NewRecorder() flash := consumeFlash(next, r) if flash == nil { t.Fatal("no flash read back") } if flash.Kind != "success" || flash.Message != "Preset created." { t.Errorf("flash = %+v", flash) } cleared := next.Result().Cookies() if len(cleared) != 1 || cleared[0].MaxAge >= 0 { t.Errorf("flash cookie was not expired: %+v", cleared) } } // The cookie is client-editable, so an arbitrary kind must not reach the // banner's class name. func TestConsumeFlashClampsKind(t *testing.T) { cases := map[string]string{ "success": "success", "error": "error", "evil\" onload=alert1": "error", } for kind, want := range cases { r := httptest.NewRequest("GET", "/", nil) r.AddCookie(&http.Cookie{Name: flashCookie, Value: url.QueryEscape(kind + "|hello")}) flash := consumeFlash(httptest.NewRecorder(), r) if flash == nil { t.Fatalf("kind %q: no flash read back", kind) } if flash.Kind != want { t.Errorf("kind %q became %q, want %q", kind, flash.Kind, want) } } } func TestConsumeFlashIgnoresMalformedValues(t *testing.T) { for _, value := range []string{"", "no-separator", "%zz"} { r := httptest.NewRequest("GET", "/", nil) r.AddCookie(&http.Cookie{Name: flashCookie, Value: value}) if flash := consumeFlash(httptest.NewRecorder(), r); flash != nil { t.Errorf("value %q produced a flash: %+v", value, flash) } } // No cookie at all. if flash := consumeFlash(httptest.NewRecorder(), httptest.NewRequest("GET", "/", nil)); flash != nil { t.Errorf("missing cookie produced a flash: %+v", flash) } } // An explicit ?sort= is remembered, and a request without one restores the last // choice — the auto-refresh reloads these pages without the query string. func TestSortFromRequest(t *testing.T) { w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/library?sort=title", nil) if got := sortFromRequest(w, r, "library_sort", "date"); got != "title" { t.Fatalf("explicit sort = %q, want title", got) } cookies := w.Result().Cookies() if len(cookies) != 1 || cookies[0].Name != "library_sort" || cookies[0].Value != "title" { t.Fatalf("sort was not remembered: %+v", cookies) } next := httptest.NewRequest("GET", "/library", nil) next.AddCookie(cookies[0]) if got := sortFromRequest(httptest.NewRecorder(), next, "library_sort", "date"); got != "title" { t.Errorf("restored sort = %q, want title", got) } // With neither query nor cookie, the caller's default applies. bare := httptest.NewRequest("GET", "/library", nil) if got := sortFromRequest(httptest.NewRecorder(), bare, "library_sort", "date"); got != "date" { t.Errorf("default sort = %q, want date", got) } } // The wildcard arrives still percent-encoded. '+' is a literal plus in a path, // so an item directory named "a+b" must round-trip. func TestNormalizeRelPath(t *testing.T) { cases := map[string]string{ "folder/item": "folder/item", "/folder/item/": "folder/item", "a%2Bb": "a+b", "a+b": "a+b", "spaced%20name": "spaced name", "%E6%97%A5%E6%9C%AC": "日本", } for raw, want := range cases { r := httptest.NewRequest("GET", "/library/item/"+raw, nil) ctx := chi.NewRouteContext() ctx.URLParams.Add("*", raw) r = r.WithContext(context.WithValue(r.Context(), chi.RouteCtxKey, ctx)) if got := normalizeRelPath(r); got != want { t.Errorf("normalizeRelPath(%q) = %q, want %q", raw, got, want) } } } func TestParseIDRejectsNonNumeric(t *testing.T) { r := httptest.NewRequest("POST", "/queue/abc/delete", nil) ctx := chi.NewRouteContext() ctx.URLParams.Add("id", "abc") r = r.WithContext(context.WithValue(r.Context(), chi.RouteCtxKey, ctx)) w := httptest.NewRecorder() if _, ok := parseID(w, r); ok { t.Error("parseID accepted a non-numeric id") } if w.Code != http.StatusBadRequest { t.Errorf("status = %d, want 400", w.Code) } } func TestFormatDuration(t *testing.T) { cases := map[int]string{ -1: "--:--", 0: "--:--", 59: "0:59", 61: "1:01", 3661: "1:01:01", } for seconds, want := range cases { if got := formatDuration(seconds); got != want { t.Errorf("formatDuration(%d) = %q, want %q", seconds, got, want) } } }