33import base64
44import json
55import logging
6+ import string
67from typing import Any
78
9+ import requests
810from django .conf import settings
911from django .db import transaction
1012from django .db .models import Q
1315from ninja import Form , Schema
1416from ninja .pagination import RouterPaginated
1517from oauth2_provider .models import AccessToken
16- import requests
1718
1819from 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 )
2522from 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.
2824from bats_ai .core .views .recording import PulseMetadataSlopesSchema
2925from 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+
7175def decode_jwt (token ):
7276 # Split the token into parts
7377 parts = token .split ("." )
@@ -89,7 +93,6 @@ def decode_jwt(token):
8993
9094def 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
167171class 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+
173183def 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 )
204252def 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):
308356def 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 )
402448def 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(
431476def 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(
454498def 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 })
472515def 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(
582629def 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
0 commit comments