Skip to content

Commit d9ba9a5

Browse files
committed
refactor: add public import support for Trakt and improve OAuth redirect handling
- Add public username import option for Trakt alongside OAuth import - Update OAuth redirect URIs to use request.build_absolute_uri() with reverse() - Split Trakt and Simkl import endpoints into separate public/private routes - Add comprehensive tests for both OAuth and public import flows - Improve documentation and type hints for import functions
1 parent cc07403 commit d9ba9a5

8 files changed

Lines changed: 229 additions & 43 deletions

File tree

src/integrations/imports/anilist.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
import requests
66
from django.apps import apps
77
from django.conf import settings
8+
from django.urls import reverse
89
from django.utils import timezone
910

1011
import app
@@ -18,8 +19,6 @@
1819

1920
def get_token(request):
2021
"""View for getting the AniList OAuth2 token."""
21-
domain = request.get_host()
22-
scheme = request.scheme
2322
code = request.GET["code"]
2423

2524
url = "https://anilist.co/api/v2/oauth/token"
@@ -29,7 +28,7 @@ def get_token(request):
2928
"client_secret": settings.ANILIST_SECRET,
3029
"code": code,
3130
"grant_type": "authorization_code",
32-
"redirect_uri": f"{scheme}://{domain}/import/anilist/private",
31+
"redirect_uri": request.build_absolute_uri(reverse("import_anilist_private")),
3332
}
3433

3534
try:

src/integrations/imports/simkl.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33

44
import requests
55
from django.conf import settings
6+
from django.urls import reverse
67
from django.utils import timezone
78
from django.utils.dateparse import parse_datetime
89

@@ -17,8 +18,6 @@
1718

1819
def get_token(request):
1920
"""View for getting the SIMKL OAuth2 token."""
20-
domain = request.get_host()
21-
scheme = request.scheme
2221
code = request.GET["code"]
2322
url = "https://api.simkl.com/oauth/token"
2423

@@ -31,7 +30,7 @@ def get_token(request):
3130
"client_secret": settings.SIMKL_SECRET,
3231
"code": code,
3332
"grant_type": "authorization_code",
34-
"redirect_uri": f"{scheme}://{domain}",
33+
"redirect_uri": request.build_absolute_uri(reverse("import_simkl_private")),
3534
}
3635

3736
try:

src/integrations/imports/trakt.py

Lines changed: 24 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44

55
import requests
66
from django.conf import settings
7+
from django.urls import reverse
78
from django.utils.dateparse import parse_datetime
89
from django_celery_beat.models import PeriodicTask
910

@@ -21,8 +22,6 @@
2122

2223
def handle_oauth_callback(request):
2324
"""View for getting the Trakt OAuth2 token."""
24-
domain = request.get_host()
25-
scheme = request.scheme
2625
code = request.GET["code"]
2726

2827
url = "https://api.trakt.tv/oauth/token"
@@ -32,7 +31,7 @@ def handle_oauth_callback(request):
3231
"client_secret": settings.TRAKT_API_SECRET,
3332
"code": code,
3433
"grant_type": "authorization_code",
35-
"redirect_uri": f"{scheme}://{domain}/import/trakt",
34+
"redirect_uri": request.build_absolute_uri(reverse("import_trakt_private")),
3635
}
3736

3837
try:
@@ -81,16 +80,18 @@ def get_username_from_oauth(access_token):
8180
return request["username"]
8281

8382

84-
def get_access_token(refresh_token):
85-
"""View for getting the Trakt OAuth2 access token."""
83+
def get_access_token(encrypted_refresh_token):
84+
"""Get access token from encrypted refresh token."""
8685
url = "https://api.trakt.tv/oauth/token"
8786

87+
decrypted_token = helpers.decrypt(encrypted_refresh_token)
88+
8889
params = {
8990
"client_id": settings.TRAKT_API,
9091
"client_secret": settings.TRAKT_API_SECRET,
91-
"refresh_token": helpers.decrypt(refresh_token),
92+
"refresh_token": decrypted_token,
9293
"grant_type": "refresh_token",
93-
"redirect_uri": f"{settings.BASE_URL}/import/trakt",
94+
"redirect_uri": f"{settings.BASE_URL}/import/trakt/private",
9495
}
9596

9697
try:
@@ -107,7 +108,7 @@ def get_access_token(refresh_token):
107108
raise
108109

109110
# refresh tokens are one time use only
110-
update_refresh_token(refresh_token, request["refresh_token"])
111+
update_refresh_token(encrypted_refresh_token, request["refresh_token"])
111112
return request["access_token"]
112113

113114

@@ -125,8 +126,19 @@ def update_refresh_token(old_token, new_token):
125126
periodic_task.save()
126127

127128

128-
def importer(token, user, mode, username=None):
129-
"""Import the user's data from Trakt using OAuth."""
129+
def importer(token, user, mode, username):
130+
"""Import the user's data from Trakt.
131+
132+
Can import using either OAuth (token provided) or public username.
133+
When using OAuth, username should be the authenticated user's username.
134+
When using public import, username is the Trakt username and token should be None.
135+
136+
Args:
137+
token (str, optional): Encrypted OAuth2 refresh token if using OAuth else None
138+
user: Django user object to import data for
139+
mode (str): Import mode ("new" or "overwrite")
140+
username (str): Trakt username to import from
141+
"""
130142
trakt_importer = TraktImporter(username, user, mode, refresh_token=token)
131143
return trakt_importer.import_data()
132144

@@ -141,7 +153,8 @@ def __init__(self, username, user, mode, refresh_token=None):
141153
username (str): Trakt username to import from
142154
user: Django user object to import data for
143155
mode (str): Import mode ("new" or "overwrite")
144-
refresh_token (str, optional): OAuth2 refresh token if using OAuth
156+
refresh_token (str, optional): Encrypted OAuth2 refresh token if
157+
using OAuth, None for public import
145158
"""
146159
self.username = username
147160
self.user = user

src/integrations/tasks.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -78,7 +78,10 @@ def import_media(importer_func, identifier, user_id, mode, oauth_username=None):
7878

7979
@shared_task(name="Import from Trakt")
8080
def import_trakt(user_id, mode, token=None, username=None):
81-
"""Celery task for importing media data from Trakt."""
81+
"""Celery task for importing media data from Trakt.
82+
83+
Can import using either OAuth (token provided) or public username.
84+
"""
8285
return import_media(trakt.importer, token, user_id, mode, username)
8386

8487

src/integrations/tests/test_imports.py

Lines changed: 101 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -614,15 +614,109 @@ def test_process_comments(self, mock_get_metadata, mock_make_request):
614614
movie_obj = trakt_importer.bulk_media[MediaTypes.MOVIE.value][0]
615615
self.assertEqual(movie_obj.notes, "Great movie!")
616616

617-
@patch("integrations.imports.trakt.TraktImporter.import_data")
618-
def test_importer_function(self, mock_import_data):
619-
"""Test the main importer function."""
620-
mock_import_data.return_value = (1, 2, 3, 4, "No warnings")
617+
@patch("integrations.imports.trakt.TraktImporter._get_paginated_data")
618+
@patch("integrations.imports.trakt.TraktImporter._make_api_request")
619+
@patch("integrations.imports.trakt.TraktImporter._get_metadata")
620+
def test_public_import_full_flow(
621+
self,
622+
mock_get_metadata,
623+
mock_make_request,
624+
mock_get_paginated,
625+
):
626+
"""Test full import flow with public username (no OAuth)."""
627+
# Mock paginated data - history returns movies, comments returns empty
628+
mock_get_paginated.side_effect = [
629+
[
630+
{
631+
"type": "movie",
632+
"movie": {"title": "Public Movie", "ids": {"tmdb": 999}},
633+
"watched_at": "2023-01-01T00:00:00.000Z",
634+
},
635+
],
636+
[], # Empty comments
637+
]
638+
639+
# Mock API requests for watchlist and ratings
640+
mock_make_request.return_value = []
641+
642+
# Mock metadata
643+
mock_get_metadata.return_value = {
644+
"title": "Public Movie",
645+
"image": "movie.jpg",
646+
}
647+
648+
# Import with no token (public)
649+
imported_counts, _ = importer(None, self.user, "new", "public_user")
650+
651+
# Verify movie was imported
652+
self.assertEqual(imported_counts[MediaTypes.MOVIE.value], 1)
653+
self.assertEqual(Movie.objects.filter(user=self.user).count(), 1)
654+
655+
@patch("integrations.imports.trakt.TraktImporter._get_paginated_data")
656+
@patch("integrations.imports.trakt.TraktImporter._make_api_request")
657+
@patch("integrations.imports.trakt.TraktImporter._get_metadata")
658+
def test_oauth_import_full_flow(
659+
self,
660+
mock_get_metadata,
661+
mock_make_request,
662+
mock_get_paginated,
663+
):
664+
"""Test full import flow with OAuth token."""
665+
# Mock paginated data - history returns movies, comments returns empty
666+
mock_get_paginated.side_effect = [
667+
[
668+
{
669+
"type": "movie",
670+
"movie": {"title": "OAuth Movie", "ids": {"tmdb": 888}},
671+
"watched_at": "2023-01-01T00:00:00.000Z",
672+
},
673+
],
674+
[], # Empty comments
675+
]
676+
677+
# Mock API requests for watchlist and ratings
678+
mock_make_request.return_value = []
679+
680+
# Mock metadata
681+
mock_get_metadata.return_value = {
682+
"title": "OAuth Movie",
683+
"image": "movie.jpg",
684+
}
685+
686+
# Import with encrypted token (OAuth)
687+
encrypted_token = helpers.encrypt("test_refresh_token")
688+
imported_counts, _ = importer(
689+
encrypted_token,
690+
self.user,
691+
"new",
692+
"oauth_user",
693+
)
694+
695+
# Verify movie was imported
696+
self.assertEqual(imported_counts[MediaTypes.MOVIE.value], 1)
697+
self.assertEqual(Movie.objects.filter(user=self.user).count(), 1)
698+
699+
def test_trakt_importer_with_refresh_token(self):
700+
"""Test TraktImporter initialization with refresh token."""
701+
encrypted_token = helpers.encrypt("test_token")
702+
importer = TraktImporter(
703+
"testuser",
704+
self.user,
705+
"new",
706+
refresh_token=encrypted_token,
707+
)
708+
709+
self.assertEqual(importer.username, "testuser")
710+
self.assertEqual(importer.refresh_token, encrypted_token)
711+
self.assertEqual(importer.mode, "new")
621712

622-
result = importer("testuser", self.user, "new")
713+
def test_trakt_importer_without_refresh_token(self):
714+
"""Test TraktImporter initialization without refresh token (public)."""
715+
importer = TraktImporter("testuser", self.user, "new", refresh_token=None)
623716

624-
# Check that the result is passed through correctly
625-
self.assertEqual(result, (1, 2, 3, 4, "No warnings"))
717+
self.assertEqual(importer.username, "testuser")
718+
self.assertIsNone(importer.refresh_token)
719+
self.assertEqual(importer.mode, "new")
626720

627721

628722
class ImportSimkl(TestCase):

src/integrations/urls.py

Lines changed: 21 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -3,19 +3,30 @@
33
from integrations import views
44

55
urlpatterns = [
6-
path("trakt-oauth", views.trakt_oauth, name="trakt_oauth"),
7-
path("import/trakt", views.import_trakt, name="import_trakt"),
8-
path("simkl-oauth", views.simkl_oauth, name="simkl_oauth"),
9-
path("import/simkl", views.import_simkl, name="import_simkl"),
6+
path("import/trakt-oauth", views.trakt_oauth, name="trakt_oauth"),
7+
path(
8+
"import/trakt/private",
9+
views.import_trakt_private,
10+
name="import_trakt_private",
11+
),
12+
path("import/trakt/public", views.import_trakt_public, name="import_trakt_public"),
13+
path("import/simkl-oauth", views.simkl_oauth, name="simkl_oauth"),
14+
path(
15+
"import/simkl_private",
16+
views.import_simkl_private,
17+
name="import_simkl_private",
18+
),
1019
path("import/mal", views.import_mal, name="import_mal"),
1120
path("import/anilist/oauth", views.anilist_oauth, name="import_anilist_oauth"),
12-
path("import/anilist/private",
13-
views.import_anilist_private,
14-
name="import_anilist_private",
21+
path(
22+
"import/anilist/private",
23+
views.import_anilist_private,
24+
name="import_anilist_private",
1525
),
16-
path("import/anilist/public",
17-
views.import_anilist_public,
18-
name="import_anilist_public",
26+
path(
27+
"import/anilist/public",
28+
views.import_anilist_public,
29+
name="import_anilist_public",
1930
),
2031
path("import/kitsu", views.import_kitsu, name="import_kitsu"),
2132
path("import/yamtrack", views.import_yamtrack, name="import_yamtrack"),

src/integrations/views.py

Lines changed: 37 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -26,7 +26,7 @@
2626
@require_POST
2727
def trakt_oauth(request):
2828
"""View for initiating Trakt OAuth2 authorization flow."""
29-
redirect_uri = request.build_absolute_uri(reverse("import_trakt"))
29+
redirect_uri = request.build_absolute_uri(reverse("import_trakt_private"))
3030
url = "https://trakt.tv/oauth/authorize"
3131
state = {
3232
"mode": request.POST["mode"],
@@ -41,8 +41,8 @@ def trakt_oauth(request):
4141

4242

4343
@require_GET
44-
def import_trakt(request):
45-
"""View for getting the Trakt OAuth2 token."""
44+
def import_trakt_private(request):
45+
"""View for handling Trakt OAuth2 callback and scheduling private import."""
4646
oauth_callback = trakt.handle_oauth_callback(request)
4747
enc_token = helpers.encrypt(oauth_callback["refresh_token"])
4848
state_token = request.GET["state"]
@@ -72,10 +72,41 @@ def import_trakt(request):
7272
return redirect("import_data")
7373

7474

75+
@require_POST
76+
def import_trakt_public(request):
77+
"""View for importing Trakt data using public username."""
78+
username = request.POST.get("user")
79+
if not username:
80+
messages.error(request, "Trakt username is required.")
81+
return redirect("import_data")
82+
83+
mode = request.POST["mode"]
84+
frequency = request.POST["frequency"]
85+
import_time = request.POST["time"]
86+
87+
if frequency == "once":
88+
tasks.import_trakt.delay(
89+
user_id=request.user.id,
90+
mode=mode,
91+
username=username,
92+
)
93+
messages.info(request, "The task to import media from Trakt has been queued.")
94+
else:
95+
helpers.create_import_schedule(
96+
username=username,
97+
request=request,
98+
mode=mode,
99+
frequency=frequency,
100+
import_time=import_time,
101+
source="Trakt",
102+
)
103+
return redirect("import_data")
104+
105+
75106
@require_POST
76107
def simkl_oauth(request):
77108
"""View for initiating the SIMKL OAuth2 authorization flow."""
78-
redirect_uri = request.build_absolute_uri(reverse("import_simkl"))
109+
redirect_uri = request.build_absolute_uri(reverse("import_simkl_private"))
79110
url = "https://simkl.com/oauth/authorize"
80111

81112
state = {
@@ -92,7 +123,7 @@ def simkl_oauth(request):
92123

93124

94125
@require_GET
95-
def import_simkl(request):
126+
def import_simkl_private(request):
96127
"""View for getting the SIMKL OAuth2 token."""
97128
oauth_callback = simkl.get_token(request)
98129
enc_token = helpers.encrypt(oauth_callback["access_token"])
@@ -167,6 +198,7 @@ def anilist_oauth(request):
167198
f"{url}?client_id={settings.ANILIST_ID}&redirect_uri={redirect_uri}&response_type=code&state={state_token}",
168199
)
169200

201+
170202
@require_GET
171203
def import_anilist_private(request):
172204
"""View for getting the AniList OAuth2 token."""

0 commit comments

Comments
 (0)