09573c27f58069b2a91d458c087c4b26485165cc

Author
TheEdgeOfRage <git@theedgeofrage.com>
Committer
TheEdgeOfRage <git@theedgeofrage.com>
Date

Message

Add support for passing URLs to addVideo

Diff

  1diff --git a/handler/videos.go b/handler/videos.go
  2index 4758b0fe364dd7cb45cae813f5a9453e9ce75d62..31031b393ad18b9ab5fd5b6471699338d0f1f35e 100644
  3--- a/handler/videos.go
  4+++ b/handler/videos.go
  5@@ -4,6 +4,7 @@ import (
  6 	"context"
  7 	"errors"
  8 	"fmt"
  9+	"net/url"
 10 	"strconv"
 11 	"strings"
 12 	"sync"
 13@@ -221,8 +222,53 @@ func (h *handler) SetVideoProgress(ctx context.Context, videoID string, progress
 14 	return video, nil
 15 }
 16 
 17-func (h *handler) AddCustomVideo(ctx context.Context, videoID string) error {
 18-	exists, err := h.db.HasVideo(ctx, videoID)
 19+type videoInput struct {
 20+	id              string
 21+	progressSeconds int
 22+}
 23+
 24+func parseVideoInput(input string) (videoInput, error) {
 25+	if !strings.Contains(input, "/") {
 26+		return videoInput{id: input}, nil
 27+	}
 28+
 29+	u, err := url.Parse(input)
 30+	if err != nil {
 31+		return videoInput{}, fmt.Errorf("invalid URL: %w", err)
 32+	}
 33+
 34+	var videoID string
 35+	switch {
 36+	case u.Host == "youtu.be":
 37+		videoID = strings.TrimPrefix(u.Path, "/")
 38+	case strings.HasPrefix(u.Path, "/live/"):
 39+		videoID = strings.TrimPrefix(u.Path, "/live/")
 40+	default:
 41+		videoID = u.Query().Get("v")
 42+	}
 43+
 44+	if videoID == "" {
 45+		return videoInput{}, fmt.Errorf("could not extract video ID from URL")
 46+	}
 47+
 48+	result := videoInput{id: videoID}
 49+	if t := u.Query().Get("t"); t != "" {
 50+		result.progressSeconds, err = strconv.Atoi(t)
 51+		if err != nil {
 52+			return videoInput{}, fmt.Errorf("invalid t parameter: %w", err)
 53+		}
 54+	}
 55+
 56+	return result, nil
 57+}
 58+
 59+func (h *handler) AddCustomVideo(ctx context.Context, rawVideoID string) error {
 60+	input, err := parseVideoInput(rawVideoID)
 61+	if err != nil {
 62+		return err
 63+	}
 64+
 65+	exists, err := h.db.HasVideo(ctx, input.id)
 66 	if err != nil {
 67 		h.log.Error("Failed to check if video already exists", "call", "db.HasVideo", "err", err)
 68 		return err
 69@@ -233,7 +279,7 @@ func (h *handler) AddCustomVideo(ctx context.Context, videoID string) error {
 70 		return nil
 71 	}
 72 
 73-	video, err := h.youTubeClient.GetVideoMetadata(ctx, videoID)
 74+	video, err := h.youTubeClient.GetVideoMetadata(ctx, input.id)
 75 	if err != nil {
 76 		h.log.Error("Failed to get video metadata", "error", err)
 77 		return err
 78@@ -259,5 +305,13 @@ func (h *handler) AddCustomVideo(ctx context.Context, videoID string) error {
 79 		}
 80 	}
 81 
 82+	if input.progressSeconds > 0 {
 83+		_, err = h.db.SetVideoProgress(ctx, input.id, input.progressSeconds)
 84+		if err != nil {
 85+			h.log.Error("Failed to set video progress", "error", err)
 86+			return err
 87+		}
 88+	}
 89+
 90 	return nil
 91 }
 92diff --git a/handler/videos_test.go b/handler/videos_test.go
 93new file mode 100644
 94index 0000000000000000000000000000000000000000..740a93c2ae1f040ff50d4ab93ea07d3a68b0070a
 95--- /dev/null
 96+++ b/handler/videos_test.go
 97@@ -0,0 +1,79 @@
 98+package handler
 99+
100+import (
101+	"testing"
102+
103+	"github.com/stretchr/testify/assert"
104+	"github.com/stretchr/testify/require"
105+)
106+
107+func TestParseVideoInput(t *testing.T) {
108+	tests := []struct {
109+		name         string
110+		input        string
111+		wantID       string
112+		wantProgress int
113+		wantErr      bool
114+	}{
115+		{
116+			name:   "plain video ID",
117+			input:  "dQw4w9WgXcQ",
118+			wantID: "dQw4w9WgXcQ",
119+		},
120+		{
121+			name:   "www.youtube.com watch URL",
122+			input:  "https://www.youtube.com/watch?v=dQw4w9WgXcQ",
123+			wantID: "dQw4w9WgXcQ",
124+		},
125+		{
126+			name:         "youtube.com watch URL with t param",
127+			input:        "https://youtube.com/watch?v=dQw4w9WgXcQ&t=9780",
128+			wantID:       "dQw4w9WgXcQ",
129+			wantProgress: 9780,
130+		},
131+		{
132+			name:   "m.youtube.com with extra params",
133+			input:  "https://m.youtube.com/watch?v=dQw4w9WgXcQ&pp=abc",
134+			wantID: "dQw4w9WgXcQ",
135+		},
136+		{
137+			name:   "youtube.com live URL",
138+			input:  "https://www.youtube.com/live/dQw4w9WgXcQ?si=abc",
139+			wantID: "dQw4w9WgXcQ",
140+		},
141+		{
142+			name:   "youtu.be short URL",
143+			input:  "https://youtu.be/dQw4w9WgXcQ",
144+			wantID: "dQw4w9WgXcQ",
145+		},
146+		{
147+			name:         "www.youtube.com watch URL with t param",
148+			input:        "https://www.youtube.com/watch?v=dQw4w9WgXcQ&t=120",
149+			wantID:       "dQw4w9WgXcQ",
150+			wantProgress: 120,
151+		},
152+		{
153+			name:    "URL missing video ID",
154+			input:   "https://www.youtube.com/watch",
155+			wantErr: true,
156+		},
157+		{
158+			name:    "invalid t param",
159+			input:   "https://www.youtube.com/watch?v=dQw4w9WgXcQ&t=abc",
160+			wantErr: true,
161+		},
162+	}
163+
164+	for _, tt := range tests {
165+		t.Run(tt.name, func(t *testing.T) {
166+			got, err := parseVideoInput(tt.input)
167+			if tt.wantErr {
168+				require.Error(t, err)
169+				return
170+			}
171+			require.NoError(t, err)
172+			assert.Equal(t, tt.wantID, got.id)
173+			assert.Equal(t, tt.wantProgress, got.progressSeconds)
174+		})
175+	}
176+}