|
3 | 3 | # Create your views here. |
4 | 4 | from django.views.decorators.cache import cache_page |
5 | 5 | from services.models import Message, Usage, FeatureUsage, Location |
6 | | -from rest_framework import response, viewsets |
| 6 | +from rest_framework import response, viewsets, status |
7 | 7 | from rest_framework.decorators import api_view |
8 | 8 | from rest_framework.permissions import IsAuthenticatedOrReadOnly, AllowAny |
9 | 9 | from services.serializer import ( |
|
15 | 15 | import django_filters |
16 | 16 | from rest_framework.reverse import reverse |
17 | 17 | from django.http import HttpResponse |
| 18 | +from django.db import connections |
| 19 | + |
18 | 20 | import json |
19 | 21 | import datetime |
20 | 22 | import hashlib |
21 | 23 | import services.plots as plotsfile |
| 24 | +from os import environ |
| 25 | +from hmac import compare_digest |
| 26 | +import logging |
| 27 | + |
| 28 | +logger = logging.getLogger(__name__) |
22 | 29 |
|
23 | 30 | OS_NAMES = ["Linux", "Windows NT", "Darwin"] |
24 | 31 | UTC = datetime.tzinfo("UTC") |
@@ -328,6 +335,71 @@ def by_root(request, format=None): |
328 | 335 | ) |
329 | 336 |
|
330 | 337 |
|
| 338 | +@api_view(("POST",)) |
| 339 | +def query(request, format=None): |
| 340 | + if not verify_token(request): |
| 341 | + logger.warning("Unauthorized query attempt") |
| 342 | + return response.Response( |
| 343 | + status=status.HTTP_401_UNAUTHORIZED, data="UNAUTHORIZED" |
| 344 | + ) |
| 345 | + |
| 346 | + param_err, sql = get_parameter(request, "sql") |
| 347 | + if param_err: |
| 348 | + logger.warning(f"Invalid query parameters: {param_err}") |
| 349 | + return response.Response( |
| 350 | + status=status.HTTP_400_BAD_REQUEST, data=f"Invalid Parameters: {param_err}" |
| 351 | + ) |
| 352 | + |
| 353 | + if not sql: |
| 354 | + logger.warning("No sql parameter provided") |
| 355 | + return response.Response( |
| 356 | + status=status.HTTP_400_BAD_REQUEST, data="No sql parameter provided" |
| 357 | + ) |
| 358 | + |
| 359 | + try: |
| 360 | + conn = connections["readonly"] |
| 361 | + with conn.cursor() as cur: |
| 362 | + cur.execute(sql) |
| 363 | + res = cur.fetchall() |
| 364 | + return response.Response(res) |
| 365 | + except Exception: |
| 366 | + logger.exception("Query execution failed") |
| 367 | + return response.Response( |
| 368 | + {"error": "Query failed"}, status=status.HTTP_400_BAD_REQUEST |
| 369 | + ) |
| 370 | + |
| 371 | + |
| 372 | +def get_bearer_token(request): |
| 373 | + """ |
| 374 | + Expect: Authorization: Bearer <token> |
| 375 | + """ |
| 376 | + auth = request.headers.get("Authorization", "") |
| 377 | + if not auth: |
| 378 | + logger.warning("No Authorization header provided") |
| 379 | + return None |
| 380 | + parts = auth.split(None, 1) # ["Bearer", "<token>"] |
| 381 | + if len(parts) != 2 or parts[0].lower() != "bearer": |
| 382 | + logger.warning("Invalid Authorization header format") |
| 383 | + return None |
| 384 | + return parts[1].strip() or None |
| 385 | + |
| 386 | + |
| 387 | +def get_parameter(request, param): |
| 388 | + val = request.POST.get(param) |
| 389 | + if val is None or val.strip() == "": |
| 390 | + return f"No {param} parameter provided", None |
| 391 | + return None, val |
| 392 | + |
| 393 | + |
| 394 | +def verify_token(request) -> bool: |
| 395 | + token = get_bearer_token(request) |
| 396 | + secret = environ.get("QUERY_SECRET_KEY", "") |
| 397 | + if not token or not secret: |
| 398 | + logger.warning("Missing token or secret") |
| 399 | + return False |
| 400 | + return compare_digest(token, secret) |
| 401 | + |
| 402 | + |
331 | 403 | class FeatureViewSet(viewsets.ModelViewSet): |
332 | 404 | """ |
333 | 405 | A viewset that provides the standard actions, |
|
0 commit comments