Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion geoapi/celery_app.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@

# Import task modules
app.conf.imports = (
"geoapi.tasks.health",
"geoapi.tasks.raster",
"geoapi.tasks.point_cloud",
"geoapi.tasks.streetview",
Expand All @@ -29,7 +30,11 @@

# Define the queues
app.conf.task_queues = {
"default": {"exchange": "default", "routing_key": "default"},
"default": {
"exchange": "default",
"routing_key": "default",
"queue_arguments": {"x-max-priority": 10},
},
"heavy": {"exchange": "heavy", "routing_key": "heavy"},
}

Expand Down
60 changes: 60 additions & 0 deletions geoapi/routes/status.py
Original file line number Diff line number Diff line change
@@ -1,17 +1,77 @@
from typing import Literal
from litestar import Controller, get, Request
from litestar.response import Response
from pydantic import BaseModel
from celery.result import AsyncResult
from geoapi.tasks.health import check_worker
from geoapi.log import logging
from geoapi.utils.decorators import not_anonymous_guard

logger = logging.getLogger(__name__)

WORKER_CHECK_TIMEOUT = (
40 # seconds; request timeout is 60s, leaving buffer for response
)


class StatusResponse(BaseModel):
status: str


class ComponentStatus(BaseModel):
status: Literal["ok", "error"]
detail: str | None = None


class WorkerStatusResponse(BaseModel):
overall: Literal["ok", "error"]
components: dict[str, ComponentStatus]


class StatusController(Controller):
path = "/status"

@get("/", tags=["status"])
async def get_status(self, request: Request) -> StatusResponse:
"""Unauthenticated liveness check for load balancer."""
return StatusResponse(status="OK")

@get("/complete", tags=["status"], guards=[not_anonymous_guard])
async def get_status_complete(
self, request: Request
) -> Response[WorkerStatusResponse]:
"""
Authenticated health check. Submits a Celery task and waits for
the result, validating worker + Redis connectivity from within the
worker network.
"""
# Submit at highest priority so health checks aren't blocked by queued default tasks
task = check_worker.apply_async(priority=10)
try:
result = AsyncResult(task.id).get(timeout=WORKER_CHECK_TIMEOUT)
except Exception as e:
logger.error(
f"Worker health check timed out or failed for user:{request.user.username}: {e}"
)
body = WorkerStatusResponse(
overall="error",
components={
"worker": ComponentStatus(
status="error", detail="Worker unavailable or timed out"
)
},
)
return Response(content=body, status_code=503)

body = WorkerStatusResponse(
overall=result["overall"],
components={
k: ComponentStatus(**v) for k, v in result["components"].items()
},
)
status_code = 200 if result["overall"] == "ok" else 503
if status_code != 200:
logger.warning(
f"Worker health check degraded/failed for user:{request.user.username} result:{result}"
)
return Response(content=body, status_code=status_code)
46 changes: 46 additions & 0 deletions geoapi/tasks/health.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
import redis
from redis.exceptions import RedisError
from sqlalchemy import text
from geoapi.celery_app import app
from geoapi.settings import settings
from geoapi.db import create_task_session
from geoapi.log import logging

logger = logging.getLogger(__name__)


def check_redis() -> dict:
try:
r = redis.Redis(host=settings.REDIS_HOST, port=settings.REDIS_PORT, db=0)
r.ping()
return {"status": "ok", "detail": None}
except RedisError as e:
logger.error(f"Worker health check: Redis problem: {e}")
return {"status": "error", "detail": str(e)}


def check_database() -> dict:
try:
with create_task_session() as session:
session.execute(text("SELECT 1"))
return {"status": "ok", "detail": None}
except Exception as e:
logger.error(f"Worker health check: DB problem: {e}")
return {"status": "error", "detail": str(e)}


def derive_overall_status(components: dict) -> str:
statuses = [c["status"] for c in components.values()]
if all(s == "ok" for s in statuses):
return "ok"
return "error"


@app.task
def check_worker() -> dict:
components = {
"redis": check_redis(),
"database": check_database(),
}
overall = derive_overall_status(components)
return {"overall": overall, "components": components}
58 changes: 58 additions & 0 deletions geoapi/tests/api_tests/test_status_routes.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,7 @@
from unittest.mock import patch, MagicMock
from celery.exceptions import TimeoutError as CeleryTimeoutError


def test_get_status_unauthorized_guest(test_client, projects_fixture):
resp = test_client.get("/status/")
data = resp.json()
Expand All @@ -10,3 +14,57 @@ def test_get_status(test_client, user1):
data = resp.json()
assert resp.status_code == 200
assert data == {"status": "OK"}


def test_get_status_complete_unauthorized(test_client):
resp = test_client.get("/status/complete")
assert resp.status_code == 401


def test_get_status_complete_ok(test_client, user1):
mock_result = {
"overall": "ok",
"components": {
"redis": {"status": "ok", "detail": None},
"database": {"status": "ok", "detail": None},
},
}
with patch("geoapi.routes.status.check_worker") as mock_task, patch(
"geoapi.routes.status.AsyncResult"
) as mock_async_result:
mock_task.delay.return_value = MagicMock(id="fake-id")
mock_async_result.return_value.get.return_value = mock_result

resp = test_client.get("/status/complete", headers={"X-Tapis-Token": user1.jwt})
assert resp.status_code == 200
assert resp.json()["overall"] == "ok"


def test_get_status_complete_error(test_client, user1):
mock_result = {
"overall": "error",
"components": {
"redis": {"status": "error", "detail": "Timeout connecting to server"},
"database": {"status": "ok", "detail": None},
},
}
with patch("geoapi.routes.status.check_worker") as mock_task, patch(
"geoapi.routes.status.AsyncResult"
) as mock_async_result:
mock_task.delay.return_value = MagicMock(id="fake-id")
mock_async_result.return_value.get.return_value = mock_result

resp = test_client.get("/status/complete", headers={"X-Tapis-Token": user1.jwt})
assert resp.status_code == 503
assert resp.json()["overall"] == "error"


def test_get_status_complete_worker_timeout(test_client, user1):
with patch("geoapi.routes.status.check_worker") as mock_task, patch(
"geoapi.routes.status.AsyncResult"
) as mock_async_result:
mock_task.delay.return_value = MagicMock(id="fake-id")
mock_async_result.return_value.get.side_effect = CeleryTimeoutError()

resp = test_client.get("/status/complete", headers={"X-Tapis-Token": user1.jwt})
assert resp.status_code == 503
Loading