diff --git a/cmd/server/main_test.go b/cmd/server/main_test.go index f065b7c..5dcd095 100644 --- a/cmd/server/main_test.go +++ b/cmd/server/main_test.go @@ -61,3 +61,15 @@ func TestNotFound(t *testing.T) { t.Fatalf("got status %d, want 404", rec.Code) } } + +func TestInvalidVideoID(t *testing.T) { + t.Parallel() + h := newTestHandler(t) + req := httptest.NewRequest(http.MethodGet, "/favicon.png", nil) + rec := httptest.NewRecorder() + h.ServeHTTP(rec, req) + + if rec.Code != http.StatusBadRequest { + t.Fatalf("got status %d, want 400", rec.Code) + } +} diff --git a/internal/handlers/handlers.go b/internal/handlers/handlers.go index cf17852..75c42a5 100644 --- a/internal/handlers/handlers.go +++ b/internal/handlers/handlers.go @@ -8,6 +8,7 @@ import ( "html/template" "log/slog" "net/http" + "regexp" "strings" yt "github.com/shanehull/yt-transcript" @@ -132,6 +133,8 @@ type pageData struct { Version string } +var videoIDRe = regexp.MustCompile(`^[A-Za-z0-9_-]{11}$`) + var indexPage = template.Must(template.New("index").Parse(indexHTML)) // Healthz returns a 200 OK status for health checks. @@ -160,8 +163,8 @@ func Index(baseURL, version string) http.Handler { func Transcript(client *yt.Client, transcriptCache *cache.Cache) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { videoID := r.PathValue("video_id") - if videoID == "" { - writeError(w, http.StatusBadRequest, "missing video_id") + if videoID == "" || !videoIDRe.MatchString(videoID) { + writeError(w, http.StatusBadRequest, "invalid video_id") return }