Skip to content

Commit cf70d8b

Browse files
committed
WIP Use keycloak to handle auth
1 parent bc03d3d commit cf70d8b

24 files changed

Lines changed: 555 additions & 136 deletions

‎bats_ai/core/views/nabat/nabat_recording.py‎

Lines changed: 79 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -3,8 +3,10 @@
33
import base64
44
import json
55
import logging
6+
import string
67
from typing import Any
78

9+
import requests
810
from django.conf import settings
911
from django.db import transaction
1012
from django.db.models import Q
@@ -13,17 +15,11 @@
1315
from ninja import Form, Schema
1416
from ninja.pagination import RouterPaginated
1517
from oauth2_provider.models import AccessToken
16-
import requests
1718

1819
from bats_ai.core.models import ProcessingTask, ProcessingTaskType, Species
19-
from bats_ai.core.models.nabat import (
20-
NABatCompressedSpectrogram,
21-
NABatPulseMetadata,
22-
NABatRecording,
23-
NABatRecordingAnnotation,
24-
)
20+
from bats_ai.core.models.nabat import (NABatCompressedSpectrogram, NABatPulseMetadata,
21+
NABatRecording, NABatRecordingAnnotation)
2522
from bats_ai.core.tasks.nabat.nabat_data_retrieval import nabat_recording_initialize
26-
2723
# Real (not TYPE_CHECKING) import: pydantic needs this at runtime to build NABatPulseMetadataSchema.
2824
from bats_ai.core.views.recording import PulseMetadataSlopesSchema
2925
from bats_ai.core.views.species import SpeciesSchema
@@ -68,6 +64,14 @@ def admin_auth(request):
6864
"""
6965

7066

67+
def get_auth_header(request: HttpRequest):
68+
auth_header = request.headers.get("Authorization")
69+
if not auth_header:
70+
return None
71+
auth_header_parts = auth_header.split(" ")
72+
return auth_header_parts[1] if len(auth_header_parts) > 1 else None
73+
74+
7175
def decode_jwt(token):
7276
# Split the token into parts
7377
parts = token.split(".")
@@ -89,7 +93,6 @@ def decode_jwt(token):
8993

9094
def get_email_if_authorized( # noqa: PLR0911
9195
request: HttpRequest,
92-
api_token: str,
9396
recording_id: int | None = None,
9497
recording_pk: int | None = None,
9598
) -> str | JsonResponse:
@@ -106,6 +109,7 @@ def get_email_if_authorized( # noqa: PLR0911
106109
if request.user and request.user.is_authenticated and request.user.is_superuser:
107110
return request.user.email or "superuser@nabat.org"
108111
# Decode JWT token
112+
api_token = get_auth_header(request)
109113
try:
110114
payload = decode_jwt(api_token)
111115
email = payload.get("email")
@@ -165,11 +169,17 @@ class NABatRecordingSchema(Schema):
165169

166170

167171
class NABatRecordingGenerateSchema(Schema):
168-
apiToken: str
169172
recordingId: int
170173
surveyEventId: int
171174

172175

176+
class NABatAuthorizationSchema(Schema):
177+
recordingId: int
178+
surveyEventId: int
179+
iss: str
180+
code: str
181+
182+
173183
def update_nabat_species(species_id: int, api_token: str, recording_id: int, survey_event_id: int):
174184
"""
175185
Update the species for a NABat recording using the NABat API.
@@ -200,6 +210,44 @@ def update_nabat_species(species_id: int, api_token: str, recording_id: int, sur
200210
return "NABat species updated successfully."
201211

202212

213+
def _reconstruct_redirect_url(survey_event_id, recording_id):
214+
url_root = settings.BATAI_WEB_URL.rstrip('/')
215+
return (
216+
f"{url_root}/nabat/auth/?recordingId={recording_id}&surveyEventId={survey_event_id}"
217+
)
218+
219+
220+
@router.post("/authorize", auth=None)
221+
def authorize_nabat_requests(
222+
request: HttpRequest,
223+
payload: Form[NABatAuthorizationSchema]
224+
):
225+
if payload.iss != settings.BATAI_NABAT_OIDC_ISSUER:
226+
return JsonResponse({"error": "Unexpected issuer"}, status=400)
227+
228+
try:
229+
redirect_url = _reconstruct_redirect_url(payload.surveyEventId, payload.recordingId)
230+
response = requests.post(
231+
f"{settings.BATAI_NABAT_OIDC_BASE_URL}/protocol/openid-connect/token",
232+
data={
233+
"grant_type": "authorization_code",
234+
"client_id": settings.BATAI_NABAT_OIDC_CLIENT_ID,
235+
"client_secret": settings.BATAI_NABAT_OIDC_CLIENT_SECRET,
236+
"redirect_uri": redirect_url,
237+
"code": payload.code,
238+
},
239+
timeout=30,
240+
)
241+
response_json = response.json()
242+
if response.status_code != 200:
243+
logger.error("Keycloak token exchange rejected: %s - %s", response.status_code, response.text)
244+
return JsonResponse({"error": "Keycloak token exchange failed."}, status=response.status_code)
245+
return JsonResponse(response_json, status=200)
246+
except Exception as e:
247+
logger.exception(e)
248+
return JsonResponse({"error": "Keycloak token exchange failed"}, status=500)
249+
250+
203251
@router.post("/", auth=None)
204252
def generate_nabat_recording( # noqa: PLR0911
205253
request: HttpRequest,
@@ -221,10 +269,11 @@ def generate_nabat_recording( # noqa: PLR0911
221269
)
222270

223271
nabat_recording = NABatRecording.objects.filter(recording_id=payload.recordingId)
272+
api_token = get_auth_header(request)
224273
if not nabat_recording.exists():
225274
# use a task to start downloading the file using the API key and generate the spectrograms
226275
task = nabat_recording_initialize.delay(
227-
payload.recordingId, payload.surveyEventId, payload.apiToken
276+
payload.recordingId, payload.surveyEventId, api_token
228277
)
229278
with transaction.atomic():
230279
ProcessingTask.objects.create(
@@ -237,9 +286,8 @@ def generate_nabat_recording( # noqa: PLR0911
237286
celery_id=task.id,
238287
)
239288
return {"taskId": task.id}
240-
# we want to check the apiToken and make sure the user has access to the file
289+
# we want to check the api token and make sure the user has access to the file
241290
# before returning it
242-
api_token = payload.apiToken
243291
recording_id = payload.recordingId
244292
headers = {"Authorization": f"Bearer {api_token}", "Content-Type": "application/json"}
245293
batch_query = QUERY % {
@@ -308,11 +356,10 @@ def get_spectrogram(request: HttpRequest, pk: int):
308356
def get_spectrogram_compressed(
309357
request: HttpRequest,
310358
pk: int,
311-
apiToken: str, # noqa: N803
312359
):
313360
nabat_recording = get_object_or_404(NABatRecording, pk=pk)
314361

315-
email_or_response = get_email_if_authorized(request, apiToken, nabat_recording.recording_id)
362+
email_or_response = get_email_if_authorized(request, nabat_recording.recording_id)
316363
if isinstance(email_or_response, JsonResponse):
317364
return email_or_response
318365

@@ -395,16 +442,14 @@ class NABatCreateRecordingAnnotationSchema(Schema):
395442
comments: str = None
396443
model: str = None
397444
confidence: float
398-
apiToken: str
399445

400446

401447
@router.get("/{nabat_recording_id}/recording-annotations", auth=admin_auth)
402448
def get_nabat_recording_annotation(
403449
request: HttpRequest,
404450
nabat_recording_id: int,
405-
apiToken: str | None = None, # noqa: N803
406451
):
407-
email_or_response = get_email_if_authorized(request, apiToken, recording_pk=nabat_recording_id)
452+
email_or_response = get_email_if_authorized(request, recording_pk=nabat_recording_id)
408453
if isinstance(email_or_response, JsonResponse):
409454
return email_or_response
410455
user_email = email_or_response # safe to use
@@ -431,9 +476,8 @@ def get_nabat_recording_annotation(
431476
def get_recording_annotation(
432477
request: HttpRequest,
433478
pk: int,
434-
apiToken: str, # noqa: N803
435479
):
436-
email_or_response = get_email_if_authorized(request, apiToken, recording_pk=pk)
480+
email_or_response = get_email_if_authorized(request, recording_pk=pk)
437481
if isinstance(email_or_response, JsonResponse):
438482
return email_or_response
439483
user_email = email_or_response # safe to use
@@ -454,9 +498,8 @@ def get_recording_annotation(
454498
def get_recording_annotation_details(
455499
request: HttpRequest,
456500
pk: int,
457-
apiToken: str, # noqa: N803
458501
):
459-
email_or_response = get_email_if_authorized(request, apiToken, recording_pk=pk)
502+
email_or_response = get_email_if_authorized(request, recording_pk=pk)
460503
if isinstance(email_or_response, JsonResponse):
461504
return email_or_response
462505
user_email = email_or_response # safe to use
@@ -471,13 +514,14 @@ def get_recording_annotation_details(
471514
@router.put("recording-annotation", auth=None, response={200: str})
472515
def create_recording_annotation(request: HttpRequest, data: NABatCreateRecordingAnnotationSchema):
473516
email_or_response = get_email_if_authorized(
474-
request, data.apiToken, recording_pk=data.recordingId
517+
request, recording_pk=data.recordingId
475518
)
476519
if isinstance(email_or_response, JsonResponse):
477520
return email_or_response
478521
user_email = email_or_response # safe to use
479522

480-
token_data = decode_jwt(data.apiToken)
523+
api_token = get_auth_header(request)
524+
token_data = decode_jwt(api_token)
481525
user_id = token_data["sub"]
482526

483527
recording = get_object_or_404(NABatRecording, pk=data.recordingId)
@@ -513,7 +557,7 @@ def update_recording_annotation(
513557
success message or an error message if the recording or species are not found.
514558
"""
515559
email_or_response = get_email_if_authorized(
516-
request, data.apiToken, recording_pk=data.recordingId
560+
request, recording_pk=data.recordingId
517561
)
518562
if isinstance(email_or_response, JsonResponse):
519563
return email_or_response
@@ -545,7 +589,7 @@ def update_nabat_recording_annotation(
545589
):
546590
"""Update an existing recording annotation in NABat."""
547591
email_or_response = get_email_if_authorized(
548-
request, data.apiToken, recording_pk=data.recordingId
592+
request, recording_pk=data.recordingId
549593
)
550594
if isinstance(email_or_response, JsonResponse):
551595
return email_or_response
@@ -569,9 +613,12 @@ def update_nabat_recording_annotation(
569613
status=400,
570614
)
571615
species_id = data.species[0]
616+
# We can pull the Authorization token out of the headers because
617+
# if it didn't exist we would have returned already
618+
api_token = get_auth_header(request)
572619
return update_nabat_species(
573620
species_id,
574-
data.apiToken,
621+
api_token,
575622
annotation.nabat_recording.recording_id,
576623
annotation.nabat_recording.survey_event_id,
577624
)
@@ -582,10 +629,9 @@ def update_nabat_recording_annotation(
582629
def delete_recording_annotation(
583630
request: HttpRequest,
584631
pk: int,
585-
apiToken: str, # noqa: N803
586632
recordingId: str, # noqa: N803
587633
):
588-
email_or_response = get_email_if_authorized(request, apiToken, recording_pk=recordingId)
634+
email_or_response = get_email_if_authorized(request, recording_pk=recordingId)
589635
if isinstance(email_or_response, JsonResponse):
590636
return email_or_response
591637
user_email = email_or_response # safe to use
@@ -646,10 +692,10 @@ def linestring_to_list(ls):
646692

647693

648694
@router.get("/{pk}/pulse_contours", auth=None)
649-
def get_pulse_contours(request: HttpRequest, pk: int, api_token: str):
695+
def get_pulse_contours(request: HttpRequest, pk: int):
650696
recording = get_object_or_404(NABatRecording, pk=pk)
651697

652-
email_or_response = get_email_if_authorized(request, api_token, recording.recording_id)
698+
email_or_response = get_email_if_authorized(request, recording.recording_id)
653699
if isinstance(email_or_response, JsonResponse):
654700
return email_or_response
655701

@@ -660,10 +706,10 @@ def get_pulse_contours(request: HttpRequest, pk: int, api_token: str):
660706

661707

662708
@router.get("/{pk}/pulse_metadata", auth=None)
663-
def get_pulse_data(request: HttpRequest, pk: int, api_token: str):
709+
def get_pulse_data(request: HttpRequest, pk: int):
664710
recording = get_object_or_404(NABatRecording, pk=pk)
665711

666-
email_or_response = get_email_if_authorized(request, api_token, recording.recording_id)
712+
email_or_response = get_email_if_authorized(request, recording.recording_id)
667713
if isinstance(email_or_response, JsonResponse):
668714
return email_or_response
669715

‎bats_ai/settings/base.py‎

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,9 +5,8 @@
55
from typing import Any
66

77
import django_stubs_ext
8-
from environ import Env
98
import osgeo
10-
9+
from environ import Env
1110
from resonant_settings.allauth import *
1211
from resonant_settings.celery import *
1312
from resonant_settings.django import *

‎bats_ai/settings/development.py‎

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,14 +1,13 @@
11
from __future__ import annotations
22

33
from django_extensions.utils import InternalIPS
4-
5-
from .base import *
6-
74
# Import these afterwards, to override
85
from resonant_settings.development.celery import *
96
from resonant_settings.development.debug_toolbar import *
107
from resonant_settings.development.minio_storage import *
118

9+
from .base import *
10+
1211
INSTALLED_APPS += [
1312
"debug_toolbar",
1413
"django_browser_reload",
@@ -57,3 +56,8 @@
5756
SHELL_PLUS_IMPORTS = [
5857
"from bats_ai.core import tasks",
5958
]
59+
60+
BATAI_NABAT_OIDC_CLIENT_ID: str = env.str("DJANGO_BATAI_NABAT_OIDC_CLIENT_ID", default="batai")
61+
BATAI_NABAT_OIDC_CLIENT_SECRET: str = env.str("DJANGO_BATAI_NABAT_OIDC_CLIENT_SECRET", default="batai-local-dev-secret")
62+
BATAI_NABAT_OIDC_ISSUER: str = env.str("DJANGO_BATAI_NABAT_OIDC_ISSUER", default="http://localhost:8081/auth/realms/NABAT")
63+
BATAI_NABAT_OIDC_BASE_URL: str = env.str("DJANGO_BATAI_NABAT_OIDC_BASE_URL", default="http://localhost:8081/auth/realms/NABAT")

‎bats_ai/settings/nabat_production.py‎

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,11 @@
11
from __future__ import annotations
22

3-
from .base import *
4-
53
# Import these afterwards, to override
64
from resonant_settings.production.https import *
75
from resonant_settings.production.s3_storage import *
86

7+
from .base import *
8+
99
SECRET_KEY: str = env.str("DJANGO_SECRET_KEY")
1010

1111
STORAGES["default"] = {
@@ -27,3 +27,9 @@
2727
FORCE_SCRIPT_NAME = _proxy_subpath
2828
# Work around https://code.djangoproject.com/ticket/36653
2929
STORAGES["staticfiles"].setdefault("OPTIONS", {})["base_url"] = f"{_proxy_subpath}/{STATIC_URL}"
30+
31+
32+
BATAI_NABAT_OIDC_CLIENT_ID: str = env.str("DJANGO_BATAI_NABAT_OIDC_CLIENT_ID", default="batai")
33+
BATAI_NABAT_OIDC_CLIENT_SECRET: str = env.str("DJANGO_BATAI_NABAT_OIDC_CLIENT_SECRET", default="batai-local-dev-secret")
34+
BATAI_NABAT_OIDC_ISSUER: str = env.str("DJANGO_BATAI_NABAT_OIDC_ISSUER", default="http://localhost:8081/auth/realms/NABAT")
35+
BATAI_NABAT_OIDC_BASE_URL: str = env.str("DJANGO_BATAI_NABAT_OIDC_BASE_URL", default="http://localhost:8081/auth/realms/NABAT")

0 commit comments

Comments
 (0)