diff --git a/.github/workflows/cd.yml b/.github/workflows/cd.yml index 6a00e08..e36fe91 100644 --- a/.github/workflows/cd.yml +++ b/.github/workflows/cd.yml @@ -28,3 +28,4 @@ jobs: docker-compose down docker-compose up -d --build docker system prune -f + sudo nginx -t && sudo systemctl reload nginx diff --git a/Dockerfile b/Dockerfile index d73af58..030bb28 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,4 +1,4 @@ -FROM python:3.11-slim +FROM python:3.13-slim COPY --from=ghcr.io/astral-sh/uv:latest /uv /bin/uv @@ -12,4 +12,4 @@ COPY . . EXPOSE 8000 -CMD ["sh", "-c", "uv run alembic upgrade head && uv run gunicorn app.main:app -w 4 -k uvicorn.workers.UvicornWorker -b 0.0.0.0:8000"] +CMD ["sh", "-c", "uv run alembic upgrade head && uv run gunicorn app.main:app -w 2 -k uvicorn.workers.UvicornWorker -b 0.0.0.0:8000"] diff --git a/app/api/router.py b/app/api/router.py new file mode 100644 index 0000000..bbcd044 --- /dev/null +++ b/app/api/router.py @@ -0,0 +1,28 @@ +from fastapi import APIRouter, Depends + +from app.api.v1 import ( + benchmark, + bin, + dashboard, + dlq, + jobs, + logs, + settings, + sse, + workers, +) +from app.core.security import verify_api_key + +router = APIRouter() +# NOTE: this is to require api key heeader. will add later +api_router = APIRouter(dependencies=[Depends(verify_api_key)]) + +router.include_router(jobs.router, prefix="/jobs", tags=["Jobs"]) +router.include_router(dlq.router, prefix="/dlq", tags=["DLQ"]) +router.include_router(bin.router, prefix="/bin", tags=["Bin"]) +router.include_router(settings.router, prefix="/settings", tags=["Settings"]) +router.include_router(workers.router, prefix="/workers", tags=["Workers"]) +router.include_router(logs.router, prefix="/logs", tags=["Logs"]) +router.include_router(benchmark.router, prefix="/benchmark", tags=["Benchmark"]) +router.include_router(dashboard.router, prefix="/dashboard", tags=["Dashboard"]) +router.include_router(sse.router, prefix="/sse", tags=["SSE"]) diff --git a/app/api/v1/benchmark.py b/app/api/v1/benchmark.py new file mode 100644 index 0000000..14c54bc --- /dev/null +++ b/app/api/v1/benchmark.py @@ -0,0 +1,28 @@ +from fastapi import APIRouter + +from app.schemas.benchmark import BenchmarkRequest, BenchmarkResult +from app.schemas.response import ApiResponse + +router = APIRouter() + + +@router.post( + "/run", + summary="Run queue algorithm benchmark", + response_model=ApiResponse[BenchmarkResult], +) +async def run_benchmark(body: BenchmarkRequest): + try: + from benchmark.runner import run_benchmark as _run + + result = await _run(n=body.n, algorithm=body.algorithm) + return ApiResponse[BenchmarkResult]( + message="Benchmark completed successfully.", + data=BenchmarkResult(**result), + ) + + except Exception as exc: + return ApiResponse[None]( + message="Benchmark failed.", + errors=[{"message": str(exc)}], + ), 500 diff --git a/app/api/v1/bin.py b/app/api/v1/bin.py new file mode 100644 index 0000000..57133e4 --- /dev/null +++ b/app/api/v1/bin.py @@ -0,0 +1,65 @@ +from uuid import UUID + +from fastapi import APIRouter + +from app.core.exceptions import FlintException +from app.dependencies import DBSession, PaginationParams +from app.schemas.job import JobResponse +from app.schemas.response import ApiResponse, Meta, error_response +from app.services import job + +router = APIRouter() + + +@router.get( + "", + summary="List bin (soft-deleted jobs)", + response_model=ApiResponse[list[JobResponse]], +) +async def list_bin(db: DBSession, page_params: PaginationParams): + page = page_params.page + limit = page_params.limit + jobs, total = await job.get_bin_jobs(page, limit, db) + return ApiResponse[list[JobResponse]]( + message="Bin retrieved successfully.", + data=[JobResponse.model_validate(j) for j in jobs], + meta=Meta(page=page, limit=limit, total=total), + ) + + +@router.patch( + "/{job_id}/restore", + summary="Restore a job from the bin", + response_model=ApiResponse[JobResponse], +) +async def restore_job(job_id: UUID, db: DBSession): + try: + job_result = await job.restore_job(job_id, db) + return ApiResponse[JobResponse]( + message="Job restored successfully.", + data=JobResponse.model_validate(job_result), + ) + except FlintException as exc: + return error_response( + message=exc.message, + errors=[{"message": exc.message}], + status_code=exc.status_code, + ) + + +@router.delete( + "/{job_id}", + summary="Permanently delete a job", + description=("Hard-deletes a job from the database."), + response_model=ApiResponse[None], +) +async def hard_delete_job(job_id: UUID, db: DBSession): + try: + await job.hard_delete_job(job_id, db) + return ApiResponse[None](message="Job permanently deleted.") + except FlintException as exc: + return error_response( + message=exc.message, + errors=[{"message": exc.message}], + status_code=exc.status_code, + ) diff --git a/app/api/v1/dashboard.py b/app/api/v1/dashboard.py new file mode 100644 index 0000000..6fba9e7 --- /dev/null +++ b/app/api/v1/dashboard.py @@ -0,0 +1,20 @@ +from fastapi import APIRouter + +from app.dependencies import DBSession +from app.schemas.response import ApiResponse +from app.services.job import get_job_counts_by_status + +router = APIRouter() + + +@router.get( + "/stats", + summary="Get dashboard job counts", + description="Returns job counts grouped by status for the dashboard.", + response_model=ApiResponse[dict], +) +async def get_dashboard_stats(db: DBSession): + counts = await get_job_counts_by_status(db) + return ApiResponse[dict]( + message="Dashboard stats retrieved successfully.", data=counts + ) diff --git a/app/api/v1/dlq.py b/app/api/v1/dlq.py new file mode 100644 index 0000000..0b72108 --- /dev/null +++ b/app/api/v1/dlq.py @@ -0,0 +1,68 @@ +from uuid import UUID + +from fastapi import APIRouter + +from app.core.exceptions import FlintException +from app.dependencies import DBSession, PaginationParams +from app.queues.heapq import HeapQueue +from app.schemas.job import JobResponse +from app.schemas.response import ApiResponse, Meta, error_response +from app.services import dlq as dlq_service + +router = APIRouter() + +_queue = HeapQueue() + + +@router.get("", summary="List DLQ jobs", response_model=ApiResponse[list[JobResponse]]) +async def list_dlq( + page_params: PaginationParams, + db: DBSession, +): + page = page_params.page + limit = page_params.limit + jobs, total = await dlq_service.get_dlq_jobs(page, limit, db) + return ApiResponse[list[JobResponse]]( + message="DLQ retrieved successfully.", + data=[JobResponse.model_validate(j) for j in jobs], + meta=Meta(page=page, limit=limit, total=total), + ) + + +@router.post( + "/{job_id}/retry", + summary="Retry a DLQ job", +) +async def retry_dlq_job(job_id: UUID, db: DBSession): + try: + job = await dlq_service.retry_dlq_job(job_id, db, _queue) + return ApiResponse[JobResponse]( + message="Job re-queued from DLQ successfully.", + data=JobResponse.model_validate(job), + ) + except FlintException as exc: + return error_response( + message=exc.message, + errors=[{"message": exc.message}], + status_code=exc.status_code, + ) + + +@router.delete( + "/{job_id}", + summary="Remove a job from DLQ (soft-delete)", + response_model=ApiResponse[None], +) +async def remove_from_dlq( + job_id: UUID, + db: DBSession, +): + try: + await dlq_service.remove_from_dlq(job_id, db) + return ApiResponse[None](message="Job removed from DLQ and moved to bin.") + except FlintException as exc: + return error_response( + message=exc.message, + errors=[{"message": exc.message}], + status_code=exc.status_code, + ) diff --git a/app/api/v1/jobs.py b/app/api/v1/jobs.py new file mode 100644 index 0000000..d29b150 --- /dev/null +++ b/app/api/v1/jobs.py @@ -0,0 +1,134 @@ +from typing import Annotated +from uuid import UUID + +from fastapi import APIRouter, Depends + +from app.core.exceptions import FlintException +from app.dependencies import DBSession +from app.schemas.job import ( + JobCreate, + JobFilterParams, + JobLogResponse, + JobResponse, +) +from app.schemas.response import ApiResponse, Meta, error_response +from app.services import job as job_service + +router = APIRouter() + + +@router.post( + "", + summary="Create a job", + response_model=ApiResponse[JobResponse], + status_code=201, +) +async def create_job( + body: JobCreate, + db: DBSession, +): + try: + job = await job_service.create_job(body, db) + return ApiResponse[JobResponse]( + message="Job created successfully.", + data=JobResponse.model_validate(job), + ) + except FlintException as exc: + return error_response( + message=exc.message, + errors=[{"message": exc.message}], + status_code=exc.status_code, + ) + except ValueError as exc: + return error_response( + message="Validation error.", + errors=[{"message": str(exc)}], + status_code=422, + ) + + +@router.get( + "", + summary="List jobs", + response_model=ApiResponse[list[JobResponse]], +) +async def list_jobs( + db: DBSession, + filters: Annotated[JobFilterParams, Depends(JobFilterParams)], +): + jobs, total = await job_service.get_jobs(filters, db) + return ApiResponse[list[JobResponse]]( + message="Jobs retrieved successfully.", + data=[JobResponse.model_validate(j) for j in jobs], + meta=Meta( + page=filters.page, + limit=filters.limit, + total=total, + ), + ) + + +@router.get( + "/{job_id}", + summary="Get job detail", + response_model=ApiResponse[JobResponse], +) +async def get_job( + job_id: UUID, + db: DBSession, +): + try: + job, dependency_ids, logs = await job_service.get_job_with_details(job_id, db) + data = JobResponse.model_validate(job) + data.dependencies = dependency_ids + data.logs = [JobLogResponse.model_validate(log) for log in logs] + return ApiResponse[JobResponse]( + message="Job retrieved successfully.", + data=data, + ) + except FlintException as exc: + return error_response( + message=exc.message, + errors=[{"message": exc.message}], + status_code=exc.status_code, + ) + + +@router.patch( + "/{job_id}/cancel", + summary="Cancel a job", + response_model=ApiResponse[JobResponse], +) +async def cancel_job( + job_id: UUID, + db: DBSession, +): + try: + job = await job_service.cancel_job(job_id, db) + return ApiResponse[JobResponse]( + message="Cancellation requested successfully.", + data=JobResponse.model_validate(job), + ) + except FlintException as exc: + return error_response( + message=exc.message, + errors=[{"message": exc.message}], + status_code=exc.status_code, + ) + + +@router.delete( + "/{job_id}", + summary="Soft-delete a job (move to bin)", + response_model=ApiResponse[None], +) +async def soft_delete_job(job_id: UUID, db: DBSession): + try: + await job_service.soft_delete_job(job_id, db) + return ApiResponse[None](message="Job moved to bin.") + except FlintException as exc: + return error_response( + message=exc.message, + errors=[{"message": exc.message}], + status_code=exc.status_code, + ) diff --git a/app/api/v1/logs.py b/app/api/v1/logs.py new file mode 100644 index 0000000..bd2aeda --- /dev/null +++ b/app/api/v1/logs.py @@ -0,0 +1,62 @@ +from typing import Annotated +from uuid import UUID + +from fastapi import APIRouter, Query +from sqlalchemy import func, select + +from app.dependencies import DBSession, PaginationParams +from app.models.job_log import JobLog +from app.schemas.log import LogEntryResponse +from app.schemas.response import ApiResponse, Meta + +router = APIRouter() + + +@router.get( + "", + summary="List job event logs", + description=( + "Returns structured log entries from the job_logs table. " + "Filter by event type or job_id. Ordered newest first." + ), +) +async def list_logs( + page_params: PaginationParams, + db: DBSession, + event: Annotated[str | None, Query()] = None, + job_id: Annotated[UUID | None, Query()] = None, +): + page = page_params.page + limit = page_params.limit + conditions = [] + if event: + conditions.append(JobLog.event == event) + if job_id: + conditions.append(JobLog.job_id == job_id) + + total_result = await db.execute( + select(func.count(JobLog.id)).where(*conditions) + if conditions + else select(func.count(JobLog.id)) + ) + total = total_result.scalar() or 0 + + offset = (page - 1) * limit + + query = ( + select(JobLog).order_by(JobLog.created_at.desc()).offset(offset).limit(limit) + ) + if conditions: + query = query.where(*conditions) + + result = await db.execute(query) + logs = result.scalars().all() + return ApiResponse[list[LogEntryResponse]]( + message="Logs retrieved successfully.", + data=[LogEntryResponse.model_validate(log) for log in logs], + meta=Meta( + page=page, + limit=limit, + total=total, + ), + ) diff --git a/app/api/v1/router.py b/app/api/v1/router.py deleted file mode 100644 index 0649154..0000000 --- a/app/api/v1/router.py +++ /dev/null @@ -1,5 +0,0 @@ -from fastapi import APIRouter, Depends - -from app.core.security import verify_api_key - -api_router = APIRouter(dependencies=[Depends(verify_api_key)]) diff --git a/app/api/v1/settings.py b/app/api/v1/settings.py new file mode 100644 index 0000000..e95b878 --- /dev/null +++ b/app/api/v1/settings.py @@ -0,0 +1,51 @@ +from fastapi import APIRouter + +from app.core.exceptions import FlintException +from app.dependencies import DBSession +from app.schemas.response import ApiResponse, error_response +from app.schemas.settings import SettingsUpdate +from app.services import settings + +router = APIRouter() + + +@router.get("", summary="Get all settings", response_model=ApiResponse[dict | None]) +async def get_settings(db: DBSession): + data = await settings.get_all_settings(db) + return ApiResponse[dict | None]( + message="Settings retrieved successfully.", + data=data, + ) + + +@router.patch( + "", + summary="Update settings", + description=( + "Update one or more settings. All fields are optional — " + "only provided keys are updated. Changes take effect immediately." + ), + response_model=ApiResponse[dict], +) +async def update_settings(body: SettingsUpdate, db: DBSession): + try: + updates = {k: v for k, v in body.model_dump().items() if v is not None} + + if not updates: + return error_response( + message="No settings provided to update.", + errors=[{"message": "Request body must include at least one setting."}], + status_code=422, + ) + + updated = await settings.update_settings(updates, db) + return ApiResponse[dict]( + message="Settings updated successfully.", + data=updated, + ) + except FlintException as exc: + return error_response( + message=exc.message, + errors=[{"message": exc.message}], + status_code=exc.status_code, + ) diff --git a/app/api/v1/sse.py b/app/api/v1/sse.py new file mode 100644 index 0000000..9e6d2c7 --- /dev/null +++ b/app/api/v1/sse.py @@ -0,0 +1,78 @@ +import asyncio + +import redis.asyncio as aioredis +from fastapi import APIRouter, Request +from fastapi.responses import StreamingResponse + +from app.core.config import settings +from app.core.logger import get_logger + +logger = get_logger(__name__) + +router = APIRouter() + +EVENTS_CHANNEL = "flint:events" + + +async def _event_generator(request: Request): + """ + Async generator that yields SSE-formatted strings. + """ + redis_client = aioredis.from_url( + settings.REDIS_URL, + decode_responses=True, + socket_timeout=None, + socket_connect_timeout=5, + ) + pubsub = redis_client.pubsub() + pubsub = redis_client.pubsub() + await pubsub.subscribe(EVENTS_CHANNEL) + + logger.info("sse_client_connected", path=str(request.url)) + + try: + async for message in pubsub.listen(): + if await request.is_disconnected(): + break + + if message["type"] != "message": + continue + + data = message.get("data", "") + if isinstance(data, bytes): + data = data.decode("utf-8") + + yield f"data: {data}\n\n" + + except asyncio.CancelledError: + pass + except Exception as exc: + logger.error("sse_stream_error", error=str(exc)) + finally: + await pubsub.unsubscribe(EVENTS_CHANNEL) + await pubsub.aclose() + logger.info("sse_client_disconnected") + + +@router.get( + "/stream", + summary="Live job event stream", + description=( + "Server-Sent Events stream. Connect with EventSource to receive " + "real-time job status updates. No API key required (EventSource " + "cannot set custom headers)." + ), + response_class=StreamingResponse, + tags=["SSE"], +) +async def sse_stream(request: Request): + return StreamingResponse( + _event_generator(request), + media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "X-Accel-Buffering": "no", + "Connection": "keep-alive", + "Access-Control-Allow-Origin": "*", + }, + ) diff --git a/app/api/v1/workers.py b/app/api/v1/workers.py new file mode 100644 index 0000000..c81b948 --- /dev/null +++ b/app/api/v1/workers.py @@ -0,0 +1,88 @@ +from typing import Annotated + +from fastapi import APIRouter, Depends +from redis.asyncio import Redis + +from app.core.exceptions import WorkerNotFoundException +from app.db.redis import get_redis +from app.schemas.response import ApiResponse + +router = APIRouter() + +WORKER_KEY_PATTERN = "flint:workers:*" +WORKER_CONTROL_CHANNEL = "flint:worker:control:{worker_id}" + + +@router.get( + "", + summary="List active workers", + response_model=ApiResponse[list[dict]], +) +async def list_workers(redis: Annotated[Redis, Depends(get_redis)]): + """ + Returns all workers currently registered in Redis. + Workers that have died without deregistering will disappear + automatically once their TTL (60s) expires. + """ + keys = await redis.keys(WORKER_KEY_PATTERN) + + workers = [] + for key in keys: + worker_id = key.replace("flint:workers:", "") # type: ignore + status = await redis.get(key) or "unknown" + ttl = await redis.ttl(key) + workers.append( + { + "worker_id": worker_id, + "status": status, + "ttl_seconds": ttl, + } + ) + return ApiResponse[list[dict]]( + message="Workers retrieved successfully.", + data=workers, + ) + + +@router.post( + "/{worker_id}/stop", + summary="Signal a worker to stop", + description=( + "Publishes a 'stop' command to the worker's control channel. " + "The worker finishes its current job then exits cleanly." + ), +) +async def stop_worker(worker_id: str, redis: Annotated[Redis, Depends(get_redis)]): + await _assert_worker_exists(worker_id, redis) + channel = WORKER_CONTROL_CHANNEL.format(worker_id=worker_id) + await redis.publish(channel, "stop") + return ApiResponse[dict]( + message=f"Stop signal sent to worker '{worker_id}'.", + data={"worker_id": worker_id, "command": "stop"}, + ) + + +@router.post( + "/{worker_id}/restart", + summary="Signal a worker to restart", + description=( + "Publishes a 'restart' command to the worker's control channel. " + "The worker will re-exec itself after finishing its current job." + ), +) +async def restart_worker(worker_id: str, redis: Annotated[Redis, Depends(get_redis)]): + await _assert_worker_exists(worker_id, redis) + channel = WORKER_CONTROL_CHANNEL.format(worker_id=worker_id) + await redis.publish(channel, "restart") + return ApiResponse[dict]( + message=f"Restart signal sent to worker '{worker_id}'.", + data={"worker_id": worker_id, "command": "restart"}, + ) + + +async def _assert_worker_exists(worker_id: str, redis: Redis) -> None: + """Raise WorkerNotFoundException if the worker is not in Redis.""" + key = f"flint:workers:{worker_id}" + exists = await redis.exists(key) + if not exists: + raise WorkerNotFoundException(worker_id) diff --git a/app/dependencies.py b/app/dependencies.py index a6ba72f..516adb5 100644 --- a/app/dependencies.py +++ b/app/dependencies.py @@ -4,5 +4,7 @@ from sqlalchemy.ext.asyncio.session import AsyncSession from app.db.session import get_db +from app.schemas.pagination import PageParams DBSession = Annotated[AsyncSession, Depends(get_db)] +PaginationParams = Annotated[PageParams, Depends(PageParams)] diff --git a/app/handlers/__init__.py b/app/handlers/__init__.py index c9cbeb2..8e5336b 100644 --- a/app/handlers/__init__.py +++ b/app/handlers/__init__.py @@ -1,8 +1,30 @@ -HANDLER_REGISTRY = {} +from app.handlers.base import BaseHandler +from app.handlers.email import EmailHandler +from app.handlers.log_processor import LogProcessorHandler +from app.handlers.webhook import WebhookHandler +from app.models.enum import JobType +HANDLER_REGISTRY: dict[str, type[BaseHandler]] = { + JobType.WEBHOOK_DELIVERY: WebhookHandler, + JobType.SEND_EMAIL: EmailHandler, + JobType.LOG_PROCESSING: LogProcessorHandler, +} -def get_handler(job_type: str): + +def get_handler(job_type: str) -> BaseHandler: handler_class = HANDLER_REGISTRY.get(job_type) if not handler_class: - raise ValueError(f"No handler registered for job type: {job_type}") + from app.core.exceptions import HandlerNotFoundException + + raise HandlerNotFoundException(job_type) return handler_class() + + +__all__ = [ + "BaseHandler", + "WebhookHandler", + "EmailHandler", + "LogProcessorHandler", + "HANDLER_REGISTRY", + "get_handler", +] diff --git a/app/main.py b/app/main.py index ba14e5c..13e4b86 100644 --- a/app/main.py +++ b/app/main.py @@ -1,8 +1,9 @@ from contextlib import asynccontextmanager from fastapi import FastAPI +from fastapi.middleware.cors import CORSMiddleware -from app.api.v1.router import api_router +from app.api.router import router from app.core.config import settings from app.core.exception_handlers import register_exception_handlers from app.core.logger import get_logger, setup_logging @@ -17,18 +18,34 @@ async def lifespan(app: FastAPI): logging.info("Starting up the application...") yield - logging.info("Shutting down the application...") await close_redis() + logging.info("Shutting down the application...") app = FastAPI( title=settings.APP_NAME, + description=( + "**Quietly igniting your payload, every job has a spark.**\n\n" + "Background job scheduler with priority queuing, DAG workflows, " + "retry logic, dead letter queue, and real-time status updates." + ), version="0.1.0", + docs_url="/api/v1/docs", + redoc_url="/api/v1/redoc", + openapi_url="/api/v1/openapi.json", lifespan=lifespan, ) +app.add_middleware( + CORSMiddleware, + allow_origins=["*"], + allow_credentials=True, + allow_methods=["*"], + allow_headers=["*"], +) + register_exception_handlers(app) -app.include_router(api_router, prefix=settings.API_V1_PREFIX) +app.include_router(router, prefix=settings.API_V1_PREFIX) @app.get("/") diff --git a/app/schemas/benchmark.py b/app/schemas/benchmark.py new file mode 100644 index 0000000..9eaa894 --- /dev/null +++ b/app/schemas/benchmark.py @@ -0,0 +1,35 @@ +from pydantic import BaseModel, Field, field_validator + + +class BenchmarkRequest(BaseModel): + n: int = Field( + default=10000, + ge=100, + le=1_000_000, + description="Number of jobs to insert and pop during the benchmark.", + ) + algorithm: str = Field( + default="both", + description="Which algorithm to benchmark: 'heap', 'timing_wheel', or 'both'.", + ) + + @field_validator("algorithm") + @classmethod + def validate_algorithm(cls, v: str) -> str: + if v not in ("heap", "timing_wheel", "both"): + return "both" + return v + + +class AlgorithmResult(BaseModel): + insert_time_ms: float + pop_time_ms: float + total_time_ms: float + + +class BenchmarkResult(BaseModel): + n: int + heap: AlgorithmResult | None = None + timing_wheel: AlgorithmResult | None = None + winner: str | None = None + notes: str diff --git a/app/schemas/job.py b/app/schemas/job.py index 3e6b682..7ef5b8e 100644 --- a/app/schemas/job.py +++ b/app/schemas/job.py @@ -125,7 +125,7 @@ class JobLogResponse(BaseModel): job_id: UUID event: str message: str - metadata_: dict[str, Any] | None = Field(None, alias="metadata") + metadata_: dict[str, Any] | None = Field(None, serialization_alias="metadata") created_at: datetime model_config = { diff --git a/app/schemas/log.py b/app/schemas/log.py new file mode 100644 index 0000000..73f75cd --- /dev/null +++ b/app/schemas/log.py @@ -0,0 +1,32 @@ +from datetime import datetime +from typing import Any +from uuid import UUID + +from pydantic import BaseModel, Field + + +class LogEntryResponse(BaseModel): + id: UUID + job_id: UUID + event: str + message: str + metadata_: dict[str, Any] | None = Field(None, serialization_alias="metadata") + created_at: datetime + + model_config = { + "from_attributes": True, + "populate_by_name": True, + } + + +class LogFilterParams(BaseModel): + page: int = Field(default=1, ge=1) + limit: int = Field(default=20, ge=1, le=100) + event: str | None = Field( + default=None, + description="Filter by event name e.g. 'job_completed'.", + ) + job_id: UUID | None = Field( + default=None, + description="Filter logs for a specific job.", + ) diff --git a/app/schemas/pagination.py b/app/schemas/pagination.py new file mode 100644 index 0000000..0242123 --- /dev/null +++ b/app/schemas/pagination.py @@ -0,0 +1,21 @@ +from fastapi import Query + + +class PageParams: + """Inject this as a dependency in any endpoint that needs pagination.""" + + def __init__( + self, + page: int = Query(1, ge=1, description="Page number, 1-indexed"), + size: int = Query(20, ge=1, le=100, description="Items per page"), + ): + self.page = page + self.size = size + + @property + def offset(self) -> int: + return (self.page - 1) * self.size + + @property + def limit(self) -> int: + return self.size diff --git a/app/schemas/response.py b/app/schemas/response.py index 73b5f06..a9c2fba 100644 --- a/app/schemas/response.py +++ b/app/schemas/response.py @@ -1,4 +1,7 @@ +from typing import Any + from pydantic import BaseModel, Field +from starlette.responses import JSONResponse class Meta(BaseModel): @@ -15,5 +18,27 @@ class ErrorDetail(BaseModel): class ApiResponse[T](BaseModel): message: str data: T | None = None - errors: list[ErrorDetail] = Field(default_factory=list) + errors: list[ErrorDetail | dict] = Field(default_factory=list) meta: Meta | None = None + + +def error_response( + message: str, + errors: list[dict[str, Any]], + status_code: int = 400, +) -> JSONResponse: + """ + Build an error API response. + + Args: + message: Human-readable summary of the error. + errors: List of error detail dicts. Each may have 'field' and 'message'. + status_code: HTTP status code. Default 400. + """ + body = { + "message": message, + "data": None, + "errors": [ErrorDetail(**e).model_dump() for e in errors], + "meta": None, + } + return JSONResponse(status_code=status_code, content=body) diff --git a/app/schemas/settings.py b/app/schemas/settings.py new file mode 100644 index 0000000..abcd8ea --- /dev/null +++ b/app/schemas/settings.py @@ -0,0 +1,78 @@ +import json + +from pydantic import BaseModel, Field, field_validator + + +class SettingResponse(BaseModel): + key: str + value: str + description: str | None = None + + model_config = {"from_attributes": True} + + +class SettingsMapResponse(BaseModel): + """Flat key->value map returned by GET /settings.""" + + dlq_threshold: str + alert_emails: str + scheduler_strategy: str + + +class SettingsUpdate(BaseModel): + """ + All fields optional — only provided keys are updated. + PATCH /api/v1/settings accepts any subset of these. + """ + + dlq_threshold: str | None = Field( + default=None, + description="Positive integer string. Default: '5'.", + ) + alert_emails: str | None = Field( + default=None, + description="JSON array string. Example: '[\"admin@example.com\"]'.", + ) + scheduler_strategy: str | None = Field( + default=None, + description="Either 'heap' or 'timing_wheel'.", + ) + + @field_validator("dlq_threshold") + @classmethod + def validate_threshold(cls, v: str | None) -> str | None: + if v is not None: + try: + val = int(v) + if val < 1: + raise ValueError + except (ValueError, TypeError) as e: + raise ValueError( + "dlq_threshold must be a positive integer string e.g. '5'." + ) from e + return v + + @field_validator("alert_emails") + @classmethod + def validate_emails(cls, v: str | None) -> str | None: + if v is not None: + try: + parsed = json.loads(v) + if not isinstance(parsed, list): + raise ValueError + for email in parsed: + if not isinstance(email, str) or "@" not in email: + raise ValueError(f"Invalid email address: '{email}'") + except (ValueError, TypeError, json.JSONDecodeError) as e: + raise ValueError( + "alert_emails must be a JSON array of email strings. " + 'Example: \'["admin@example.com", "oncall@example.com"]\'' + ) from e + return v + + @field_validator("scheduler_strategy") + @classmethod + def validate_strategy(cls, v: str | None) -> str | None: + if v is not None and v not in ("heap", "timing_wheel"): + raise ValueError("scheduler_strategy must be 'heap' or 'timing_wheel'.") + return v diff --git a/app/services/dlq.py b/app/services/dlq.py index b8baa20..8df85fc 100644 --- a/app/services/dlq.py +++ b/app/services/dlq.py @@ -192,12 +192,10 @@ async def retry_dlq_job( logger.info("dlq_job_retried", job_id=str(job_id)) - # Cascade reset downstream auto-cancelled dependents from app.services import dag await dag.on_dag_root_retried(job_id, db) - # Re-evaluate DAG: push to queue only if all deps are met has_unmet = await dag.has_unmet_dependencies(job_id, db) if not has_unmet: await queue.push( @@ -206,14 +204,14 @@ async def retry_dlq_job( scheduled_at=job.scheduled_at.timestamp(), created_at=job.created_at.timestamp(), ) - # Also sync to Redis sorted set - from app.services.job_service import _sync_job_to_redis + from app.services.job import _sync_job_to_redis await _sync_job_to_redis(str(job_id), float(job.priority)) await db.commit() refreshed = await db.get(Job, job_id) + assert refreshed is not None return refreshed diff --git a/benchmark/__init__.py b/benchmark/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/benchmark/runner.py b/benchmark/runner.py new file mode 100644 index 0000000..89ed486 --- /dev/null +++ b/benchmark/runner.py @@ -0,0 +1,153 @@ +import argparse +import asyncio +import random +import time + +from app.queues.heapq import HeapQueue +from app.queues.timing_wheel import TimingWheel + + +async def run_benchmark( + n: int = 10000, + algorithm: str = "both", +) -> dict: + """ + Benchmark HeapQueue and/or TimingWheel. + + Args: + n: Number of jobs to insert and pop. + algorithm: 'heap', 'timing_wheel', or 'both'. + + Returns: + Dict with timing results and winner. + """ + now = time.time() + jobs = [ + ( + str(i), + float(random.randint(1, 3)), + now + random.uniform(0, 3600), + now - random.uniform(0, 600), + ) + for i in range(n) + ] + + results: dict = {"n": n} + + # ========================= + # Heap benchmark + # ========================= + if algorithm in ("heap", "both"): + heap = HeapQueue() + + t0 = time.perf_counter() + for job_id, ep, sa, ca in jobs: + await heap.push(job_id, ep, sa, ca) + insert_time = time.perf_counter() - t0 + + t0 = time.perf_counter() + while await heap.size() > 0: + await heap.pop() + pop_time = time.perf_counter() - t0 + + results["heap"] = { + "insert_time_ms": round(insert_time * 1000, 2), + "pop_time_ms": round(pop_time * 1000, 2), + "total_time_ms": round((insert_time + pop_time) * 1000, 2), + } + + # ========================= + # Timing wheel benchmark + # ========================= + if algorithm in ("timing_wheel", "both"): + wheel = TimingWheel() + + # Insert + t0 = time.perf_counter() + for job_id, ep, sa, ca in jobs: + await wheel.push(job_id, ep, sa, ca) + insert_time = time.perf_counter() - t0 + + # Drain — tick until empty + t0 = time.perf_counter() + while await wheel.size() > 0: + await wheel.tick() + pop_time = time.perf_counter() - t0 + + results["timing_wheel"] = { + "insert_time_ms": round(insert_time * 1000, 2), + "pop_time_ms": round(pop_time * 1000, 2), + "total_time_ms": round((insert_time + pop_time) * 1000, 2), + } + + # Winner + if algorithm == "both" and "heap" in results and "timing_wheel" in results: + heap_total = results["heap"]["total_time_ms"] + tw_total = results["timing_wheel"]["total_time_ms"] + results["winner"] = "timing_wheel" if tw_total < heap_total else "heap" + results["notes"] = ( + f"At n={n}: timing_wheel total={tw_total}ms, heap total={heap_total}ms. " + "Timing wheel wins on raw insert/pop throughput (O(1) vs O(log n)). " + "Heap wins on priority ordering correctness and re-scoring efficiency " + "(aging process). For Flint's workload — priority + starvation prevention " + "— heap is the correct primary algorithm." + ) + elif algorithm == "heap": + results["winner"] = "heap" + results["notes"] = "Only heap was benchmarked." + else: + results["winner"] = "timing_wheel" + results["notes"] = "Only timing wheel was benchmarked." + + return results + + +def _print_results(results: dict) -> None: + n = results["n"] + print(f"\n{'=' * 55}") + print(f" FLINT BENCHMARK RESULTS (n={n:,})") + print(f"{'=' * 55}") + + for algo, label in [("heap", "HEAP"), ("timing_wheel", "TIMING WHEEL")]: + if algo in results: + r = results[algo] + print(f"\n {label}") + print(f" Insert : {r['insert_time_ms']:>10.2f} ms") + print(f" Pop : {r['pop_time_ms']:>10.2f} ms") + print(f" Total : {r['total_time_ms']:>10.2f} ms") + + if "winner" in results: + print(f"\n Winner (total time): {results['winner'].upper().replace('_', ' ')}") + print("\n Analysis:") + words = results["notes"].split() + line = " " + for word in words: + if len(line) + len(word) > 72: + print(line) + line = " " + word + " " + else: + line += word + " " + if line.strip(): + print(line) + + print(f"\n{'=' * 55}\n") + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description="Flint — Queue Algorithm Benchmark") + parser.add_argument( + "--n", + type=int, + default=10000, + help="Number of jobs to benchmark (default: 10000)", + ) + parser.add_argument( + "--algorithm", + choices=["heap", "timing_wheel", "both"], + default="both", + help="Algorithm to benchmark (default: both)", + ) + args = parser.parse_args() + + results = asyncio.run(run_benchmark(n=args.n, algorithm=args.algorithm)) + _print_results(results) diff --git a/docker-compose.yml b/docker-compose.yml index 788896e..f597285 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -15,17 +15,83 @@ services: retries: 5 restart: unless-stopped + redis: + image: redis:7-alpine + container_name: flint-redis + volumes: + - redis_data:/data + healthcheck: + test: ["CMD", "redis-cli", "ping"] + interval: 5s + timeout: 5s + retries: 5 + restart: unless-stopped + api: build: . container_name: flint-api ports: - "127.0.0.1:8000:8000" + env_file: .env environment: DATABASE_URL: postgresql+asyncpg://${POSTGRES_USER:-flint}:${POSTGRES_PASSWORD}@postgres:5432/${POSTGRES_DB:-flint} + REDIS_URL: redis://redis:6379/0 depends_on: postgres: condition: service_healthy + redis: + condition: service_healthy restart: unless-stopped + worker-1: + build: . + command: uv run python -m worker.worker + env_file: .env + environment: + WORKER_ID: worker-1 + DATABASE_URL: postgresql+asyncpg://${POSTGRES_USER:-flint}:${POSTGRES_PASSWORD}@postgres:5432/${POSTGRES_DB:-flint} + REDIS_URL: redis://redis:6379/0 + depends_on: + postgres: + condition: service_healthy + redis: + condition: service_healthy + volumes: + - ./logs:/app/logs + + worker-2: + build: . + command: uv run python -m worker.worker + env_file: .env + environment: + WORKER_ID: worker-2 + DATABASE_URL: postgresql+asyncpg://${POSTGRES_USER:-flint}:${POSTGRES_PASSWORD}@postgres:5432/${POSTGRES_DB:-flint} + REDIS_URL: redis://redis:6379/0 + depends_on: + postgres: + condition: service_healthy + redis: + condition: service_healthy + volumes: + - ./logs:/app/logs + + scheduler: + build: . + command: uv run python -m scheduler.scheduler + env_file: .env + environment: + DATABASE_URL: postgresql+asyncpg://${POSTGRES_USER:-flint}:${POSTGRES_PASSWORD}@postgres:5432/${POSTGRES_DB:-flint} + REDIS_URL: redis://redis:6379/0 + depends_on: + postgres: + condition: service_healthy + redis: + condition: service_healthy + volumes: + - ./logs:/app/logs + + + volumes: postgres_data: + redis_data: diff --git a/pyproject.toml b/pyproject.toml index 784aa60..68e830d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -5,6 +5,7 @@ description = "Quietly igniting your payload, every job has a spark" readme = "README.md" requires-python = ">=3.13" dependencies = [ + "aiosmtplib>=5.1.1", "alembic>=1.18.4", "asyncpg>=0.31.0", "fastapi[standard]>=0.136.3", @@ -12,6 +13,7 @@ dependencies = [ "httpx2>=2.3.0", "pydantic-settings>=2.14.1", "pytest>=9.0.3", + "python-decouple>=3.8", "redis>=8.0.0", "sqlalchemy[asyncio]>=2.0.50", "structlog>=26.1.0", diff --git a/pytest.ini b/pytest.ini new file mode 100644 index 0000000..e30d6b5 --- /dev/null +++ b/pytest.ini @@ -0,0 +1,9 @@ +[pytest] +asyncio_mode = auto +testpaths = tests +python_files = test_*.py +python_classes = Test* +python_functions = test_* +filterwarnings = + ignore::DeprecationWarning + ignore::pytest.PytestUnraisableExceptionWarning diff --git a/scheduler/scheduler.py b/scheduler/scheduler.py new file mode 100644 index 0000000..dd7e543 --- /dev/null +++ b/scheduler/scheduler.py @@ -0,0 +1,235 @@ +import asyncio +import signal +from datetime import UTC, datetime + +import redis.asyncio as aioredis +from sqlalchemy import func, select + +from app.core.config import settings +from app.core.logger import get_logger, setup_logging +from app.db.session import AsyncSessionLocal +from app.models.job import Job, JobStatus +from app.models.job_depedencies import JobDependency +from app.queues.heapq import HeapQueue +from worker.aging import AgingProcess + +setup_logging() +logger = get_logger(__name__) + +QUEUE_KEY = "flint:queue" +WORKER_KEY_PATTERN = "flint:worker:*" +DEAD_WORKER_CLEANUP_INTERVAL = 60 # seconds + + +class FlintScheduler: + def __init__(self) -> None: + self.running = False + self._redis: aioredis.Redis + self._queue = HeapQueue() + + async def start(self) -> None: + """Start all scheduler loops concurrently.""" + self.running = True + + self._redis = aioredis.from_url( + settings.REDIS_URL, + decode_responses=True, + encoding="utf-8", + ) + + loop = asyncio.get_event_loop() + for sig in (signal.SIGTERM, signal.SIGINT): + loop.add_signal_handler(sig, self._handle_shutdown) + + logger.info("scheduler_started") + + await asyncio.gather( + self._due_job_loop(), + self._aging_loop(), + self._dead_worker_cleanup_loop(), + ) + + async def _due_job_loop(self) -> None: + """ + Every SCHEDULER_POLL_INTERVAL seconds: + Find all pending jobs that are due and eligible, push to Redis queue. + """ + while self.running: + try: + await self._push_due_jobs() + except asyncio.CancelledError: + break + except Exception as exc: + logger.error("scheduler_due_job_error", error=str(exc)) + await asyncio.sleep(settings.SCHEDULER_POLL_INTERVAL) + + async def _push_due_jobs(self) -> None: + """ + Query for eligible jobs and push them to the Redis sorted set. + """ + async with AsyncSessionLocal() as session: + now = datetime.now(UTC) + + result = await session.execute( + select(Job) + .where( + Job.status == JobStatus.PENDING, + Job.scheduled_at <= now, + Job.deleted_at.is_(None), + Job.is_dlq.is_(False), + Job.worker_id.is_(None), + ) + .order_by( + Job.effective_priority.asc(), + Job.scheduled_at.asc(), + Job.created_at.asc(), + ) + .limit(100) + ) + candidate_jobs = result.scalars().all() + + if not candidate_jobs: + return + + pushed = 0 + for job in candidate_jobs: + score = await self._redis.zscore(QUEUE_KEY, str(job.id)) + if score is not None: + continue + + unmet = await self._has_unmet_dependencies(job.id, session) + if unmet: + continue + + await self._redis.zadd( + QUEUE_KEY, + {str(job.id): job.effective_priority}, + ) + pushed += 1 + + if pushed > 0: + logger.info( + "jobs_pushed_to_queue", + count=pushed, + timestamp=now.isoformat(), + ) + + async def _has_unmet_dependencies( + self, + job_id, + session, + ) -> bool: + """Return True if the job has any dependency that is not completed.""" + result = await session.execute( + select(func.count(JobDependency.id)) + .join(Job, Job.id == JobDependency.depends_on_id) + .where( + JobDependency.job_id == job_id, + Job.status != JobStatus.COMPLETED, + ) + ) + return (result.scalar() or 0) > 0 + + async def _aging_loop(self) -> None: + """ + Every AGING_INTERVAL seconds: run the starvation prevention + aging process to decrement effective_priority on waiting jobs. + """ + while self.running: + try: + async with AsyncSessionLocal() as session: + aging = AgingProcess() + await aging.run(session, self._queue) + except asyncio.CancelledError: + break + except Exception as exc: + logger.error("scheduler_aging_error", error=str(exc)) + await asyncio.sleep(settings.AGING_INTERVAL) + + async def _dead_worker_cleanup_loop(self) -> None: + """ + Every 60 seconds: find jobs stuck in 'processing' state whose + worker heartbeat has expired in Redis. Reset them to 'pending' + so they get picked up again. + """ + while self.running: + try: + await self._cleanup_dead_worker_jobs() + except asyncio.CancelledError: + break + except Exception as exc: + logger.error("scheduler_cleanup_error", error=str(exc)) + await asyncio.sleep(DEAD_WORKER_CLEANUP_INTERVAL) + + async def _cleanup_dead_worker_jobs(self) -> None: + """ + Find active worker IDs from Redis. Any job with a worker_id + not in that set has been abandoned by a dead worker and should + be reset to pending. + """ + worker_keys = await self._redis.keys(WORKER_KEY_PATTERN) + active_worker_ids = {key.replace("flint:worker:", "") for key in worker_keys} # type: ignore + + async with AsyncSessionLocal() as session: + result = await session.execute( + select(Job).where( + Job.status == JobStatus.PROCESSING, + Job.deleted_at.is_(None), + Job.worker_id.isnot(None), + ) + ) + processing_jobs = result.scalars().all() + + reset_count = 0 + for job in processing_jobs: + if job.worker_id not in active_worker_ids: + from sqlalchemy import update + + await session.execute( + update(Job) + .where(Job.id == job.id) + .values( + status=JobStatus.PENDING, + worker_id=None, + updated_at=func.now(), + ) + ) + + from app.models.job_log import JobLog, LogEvent + + log = JobLog( + job_id=job.id, + event=LogEvent.JOB_CREATED, + message=( + f"Job reset to pending: worker '{job.worker_id}' " + f"is no longer active (heartbeat expired)." + ), + metadata_={ + "dead_worker_id": job.worker_id, + "reason": "dead_worker_recovery", + }, + ) + session.add(log) + reset_count += 1 + + logger.warning( + "dead_worker_job_reset", + job_id=str(job.id), + dead_worker_id=job.worker_id, + ) + + if reset_count > 0: + await session.commit() + logger.info( + "dead_worker_cleanup_complete", + reset_count=reset_count, + ) + + def _handle_shutdown(self) -> None: + logger.info("scheduler_shutdown_signal") + self.running = False + + +if __name__ == "__main__": + scheduler = FlintScheduler() + asyncio.run(scheduler.start()) diff --git a/uv.lock b/uv.lock index 3b07718..f7c4812 100644 --- a/uv.lock +++ b/uv.lock @@ -6,6 +6,15 @@ resolution-markers = [ "python_full_version < '3.14'", ] +[[package]] +name = "aiosmtplib" +version = "5.1.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/39/ba/34f2fef90d13e21ae3f1b360da98d825c40832bb232613513be92457ff65/aiosmtplib-5.1.1.tar.gz", hash = "sha256:d9a35e9d170bc1a9f66e2fdfe7fd212f7eebb8c1581c621f79395d0bcaba7a68", size = 68123, upload-time = "2026-05-31T17:25:36.298Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/56/97/d1030d897e96c79cf0682ff93c11a2118085b3af4c27993675eda9e55da3/aiosmtplib-5.1.1-py3-none-any.whl", hash = "sha256:9d384f0c3d8906f745c1cf6819f073145bb2de8b10407905f5e2ee3389bfe6c7", size = 27937, upload-time = "2026-05-31T17:25:35.283Z" }, +] + [[package]] name = "alembic" version = "1.18.4" @@ -273,6 +282,7 @@ name = "flint" version = "0.1.0" source = { virtual = "." } dependencies = [ + { name = "aiosmtplib" }, { name = "alembic" }, { name = "asyncpg" }, { name = "fastapi", extra = ["standard"] }, @@ -280,6 +290,7 @@ dependencies = [ { name = "httpx2" }, { name = "pydantic-settings" }, { name = "pytest" }, + { name = "python-decouple" }, { name = "redis" }, { name = "sqlalchemy", extra = ["asyncio"] }, { name = "structlog" }, @@ -295,6 +306,7 @@ dev = [ [package.metadata] requires-dist = [ + { name = "aiosmtplib", specifier = ">=5.1.1" }, { name = "alembic", specifier = ">=1.18.4" }, { name = "asyncpg", specifier = ">=0.31.0" }, { name = "fastapi", extras = ["standard"], specifier = ">=0.136.3" }, @@ -302,6 +314,7 @@ requires-dist = [ { name = "httpx2", specifier = ">=2.3.0" }, { name = "pydantic-settings", specifier = ">=2.14.1" }, { name = "pytest", specifier = ">=9.0.3" }, + { name = "python-decouple", specifier = ">=3.8" }, { name = "redis", specifier = ">=8.0.0" }, { name = "sqlalchemy", extras = ["asyncio"], specifier = ">=2.0.50" }, { name = "structlog", specifier = ">=26.1.0" }, @@ -763,6 +776,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/d4/24/a372aaf5c9b7208e7112038812994107bc65a84cd00e0354a88c2c77a617/pytest-9.0.3-py3-none-any.whl", hash = "sha256:2c5efc453d45394fdd706ade797c0a81091eccd1d6e4bccfcd476e2b8e0ab5d9", size = 375249, upload-time = "2026-04-07T17:16:16.13Z" }, ] +[[package]] +name = "python-decouple" +version = "3.8" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/e1/97/373dcd5844ec0ea5893e13c39a2c67e7537987ad8de3842fe078db4582fa/python-decouple-3.8.tar.gz", hash = "sha256:ba6e2657d4f376ecc46f77a3a615e058d93ba5e465c01bbe57289bfb7cce680f", size = 9612, upload-time = "2023-03-01T19:38:38.143Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/a2/d4/9193206c4563ec771faf2ccf54815ca7918529fe81f6adb22ee6d0e06622/python_decouple-3.8-py3-none-any.whl", hash = "sha256:d0d45340815b25f4de59c974b855bb38d03151d81b037d9e3f463b0c9f8cbd66", size = 9947, upload-time = "2023-03-01T19:38:36.015Z" }, +] + [[package]] name = "python-dotenv" version = "1.2.2" diff --git a/worker/aging.py b/worker/aging.py new file mode 100644 index 0000000..e31c953 --- /dev/null +++ b/worker/aging.py @@ -0,0 +1,156 @@ +from datetime import UTC, datetime, timedelta +from typing import TYPE_CHECKING + +from sqlalchemy import func, select, update +from sqlalchemy.ext.asyncio import AsyncSession + +from app.core.config import settings +from app.core.logger import get_logger +from app.models.job import Job, JobPriority, JobStatus + +if TYPE_CHECKING: + from app.queues.base import BaseQueue + + +logger = get_logger(__name__) + +_PRIORITY_FLOOR = 1.0 + + +class AgingProcess: + """ + Decrements effective_priority on long-waiting pending jobs. + """ + + async def run( + self, + session: AsyncSession, + queue: "BaseQueue", + ) -> dict: + """ + Execute one aging cycle. + """ + now = datetime.now(UTC) + + medium_cutoff = now - timedelta(seconds=settings.MEDIUM_PRIORITY_AGE_THRESHOLD) + low_cutoff = now - timedelta(seconds=settings.LOW_PRIORITY_AGE_THRESHOLD) + + medium_count = await self._age_priority( + priority=JobPriority.MEDIUM, + cutoff=medium_cutoff, + db=session, + ) + + low_count = await self._age_priority( + priority=JobPriority.LOW, + cutoff=low_cutoff, + db=session, + ) + + await session.commit() + + if medium_count > 0 or low_count > 0: + await self._sync_queue(session, queue) + + summary = { + "medium_jobs_aged": medium_count, + "low_jobs_aged": low_count, + "total_aged": medium_count + low_count, + "timestamp": now.isoformat(), + } + + logger.info( + "aging_complete", + **summary, + ) + + return summary + + async def _age_priority( + self, + priority: int, + cutoff: datetime, + db: AsyncSession, + ) -> int: + """ + Decrement effective_priority for all pending jobs of a given + priority level that have been waiting since before the cutoff. + + Uses GREATEST() to floor at _PRIORITY_FLOOR (1.0). + Returns the number of rows updated. + """ + result = await db.execute( + update(Job) + .where( + Job.status == JobStatus.PENDING, + Job.priority == priority, + Job.effective_priority > _PRIORITY_FLOOR, + Job.created_at <= cutoff, + Job.deleted_at.is_(None), + Job.is_dlq.is_(False), + ) + .values( + effective_priority=func.greatest( + _PRIORITY_FLOOR, + Job.effective_priority - settings.AGING_DECREMENT, + ), + updated_at=func.now(), + ) + .returning(Job.id) + ) + updated_ids = result.fetchall() + count = len(updated_ids) + + if count > 0: + logger.info( + "jobs_aged", + priority=priority, + count=count, + decrement=settings.AGING_DECREMENT, + ) + + return count + + async def _sync_queue( + self, + db: AsyncSession, + queue: "BaseQueue", + ) -> None: + """ + After aging, fetch all updated pending jobs and re-push them + to the queue with their new effective_priority scores. + """ + result = await db.execute( + select(Job).where( + Job.status == JobStatus.PENDING, + Job.deleted_at.is_(None), + Job.is_dlq.is_(False), + # Only re-sync jobs that have been aged (below their raw priority) + Job.effective_priority + < Job.priority.cast(type_=type(Job.effective_priority.type)), + ) + ) + aged_jobs = result.scalars().all() + + for job in aged_jobs: + try: + await queue.update_priority( + job_id=str(job.id), + new_priority=job.effective_priority, + scheduled_at=job.scheduled_at.timestamp(), + created_at=job.created_at.timestamp(), + ) + + # Also update Redis sorted set score + from app.services.job import _sync_job_to_redis + + await _sync_job_to_redis( + str(job.id), + job.effective_priority, + ) + except Exception as exc: + logger.error( + "aging_sync_error", + job_id=str(job.id), + error=str(exc), + ) diff --git a/worker/processor.py b/worker/processor.py new file mode 100644 index 0000000..16d1c3f --- /dev/null +++ b/worker/processor.py @@ -0,0 +1,475 @@ +import asyncio +import math +import random +import uuid +from datetime import UTC, datetime, timedelta +from typing import TYPE_CHECKING + +from sqlalchemy import func, select, update +from sqlalchemy.ext.asyncio import AsyncSession + +from app.core.logger import get_logger +from app.db.session import AsyncSessionLocal +from app.handlers import get_handler +from app.models.job import Job, JobStatus +from app.models.job_log import JobLog, LogEvent +from app.services import dag as dag_service +from app.services.dlq import send_to_dlq +from app.services.job import ( + _publish_sse_event, + _sync_job_to_redis, +) + +if TYPE_CHECKING: + from app.queues.base import BaseQueue + +logger = get_logger(__name__) + + +def calculate_next_retry_delay(attempt: int) -> float: + """ + Exponential backoff with jitter. + """ + base = math.pow(5, attempt - 1) + jitter = base * random.uniform(0.5, 1.5) + return round(jitter, 2) + + +class JobProcessor: + def __init__(self, worker_id: str, queue: "BaseQueue") -> None: + self.worker_id = worker_id + self.queue = queue + + async def process(self, job_id: str) -> None: + """ + Full job processing flow for a single job. + """ + async with AsyncSessionLocal() as db: + job = await self._load_job(job_id, db) + if not job: + return + + claimed = await self._claim_job(job.id, db) + if not claimed: + logger.info( + "job_claim_failed", + job_id=job_id, + worker_id=self.worker_id, + reason="already_claimed", + ) + return + + await self._log_event( + job_id=job.id, + event=LogEvent.JOB_STARTED, + message=f"Job started by worker {self.worker_id}.", + metadata={"worker_id": self.worker_id}, + db=db, + ) + await _publish_sse_event( + { + "job_id": job_id, + "status": JobStatus.PROCESSING, + "worker_id": self.worker_id, + } + ) + + logger.info( + "job_started", + job_id=job_id, + type=job.type, + worker_id=self.worker_id, + retry_count=job.retry_count, + ) + + if await self._is_cancellation_requested(job.id, db): + await self._mark_cancelled( + job.id, + db, + reason="cancellation_requested_before_execute", + ) + return + + start_time = datetime.now(UTC) + try: + handler = get_handler(job.type) + result = await handler.execute(job.payload) + except Exception as exc: + elapsed_ms = self._elapsed_ms(start_time) + logger.warning( + "job_execution_failed", + job_id=job_id, + type=job.type, + error=str(exc), + elapsed_ms=elapsed_ms, + retry_count=job.retry_count, + ) + await self._handle_failure(job, exc, db) + return + + elapsed_ms = self._elapsed_ms(start_time) + + if await self._is_cancellation_requested(job.id, db): + await self._mark_cancelled( + job.id, + db, + reason="cancellation_requested_after_execute", + ) + return + + await db.execute( + update(Job) + .where(Job.id == job.id) + .values( + status=JobStatus.COMPLETED, + completed_at=func.now(), + worker_id=None, + updated_at=func.now(), + ) + ) + + await self._log_event( + job_id=job.id, + event=LogEvent.JOB_COMPLETED, + message=( + f"Job completed successfully by worker {self.worker_id} " + f"in {elapsed_ms}ms." + ), + metadata={ + "worker_id": self.worker_id, + "duration_ms": elapsed_ms, + "result": result, + }, + db=db, + ) + await db.commit() + + logger.info( + "job_completed", + job_id=job_id, + type=job.type, + worker_id=self.worker_id, + duration_ms=elapsed_ms, + ) + + await _publish_sse_event( + { + "job_id": job_id, + "status": JobStatus.COMPLETED, + "worker_id": self.worker_id, + "duration_ms": elapsed_ms, + } + ) + + await dag_service.on_job_completed(job.id, db, self.queue) + + if job.interval_seconds: + await self._handle_recurrence(job, db) + + async def _claim_job( + self, + job_id: uuid.UUID, + db: AsyncSession, + ) -> bool: + """ + Atomic claim via single UPDATE with WHERE guard. + + Only succeeds if status='pending' AND worker_id IS NULL. + If two workers race, exactly one gets the row back. + PostgreSQL row-level locking guarantees atomicity. + + Returns True if this worker successfully claimed the job. + """ + result = await db.execute( + update(Job) + .where( + Job.id == job_id, + Job.status == JobStatus.PENDING, + Job.worker_id.is_(None), + Job.deleted_at.is_(None), + ) + .values( + status=JobStatus.PROCESSING, + worker_id=self.worker_id, + started_at=func.now(), + updated_at=func.now(), + ) + .returning(Job.id) + ) + await db.commit() + return result.scalar_one_or_none() is not None + + async def _handle_failure( + self, + job: Job, + error: Exception, + db: AsyncSession, + ) -> None: + """ + Handle a job execution failure. + """ + new_retry_count = job.retry_count + 1 + error_str = str(error) + + if new_retry_count <= job.max_retries: + delay = calculate_next_retry_delay(new_retry_count) + next_retry_at = datetime.now(UTC) + timedelta(seconds=delay) + + await db.execute( + update(Job) + .where(Job.id == job.id) + .values( + status=JobStatus.PENDING, + retry_count=new_retry_count, + next_retry_at=next_retry_at, + last_error=error_str[:1000], + worker_id=None, + updated_at=func.now(), + ) + ) + + await self._log_event( + job_id=job.id, + event=LogEvent.JOB_RETRY_ATTEMPTED, + message=( + f"Job failed on attempt {new_retry_count}/{job.max_retries}. " + f"Retrying in {delay:.1f}s. Error: {error_str[:200]}" + ), + metadata={ + "attempt": new_retry_count, + "max_retries": job.max_retries, + "delay_seconds": delay, + "error": error_str[:500], + "next_retry_at": next_retry_at.isoformat(), + }, + db=db, + ) + await db.commit() + + logger.warning( + "job_retry_attempted", + job_id=str(job.id), + attempt=new_retry_count, + max_retries=job.max_retries, + delay_seconds=delay, + error=error_str[:200], + ) + + await _publish_sse_event( + { + "job_id": str(job.id), + "status": JobStatus.PENDING, + "retry_count": new_retry_count, + } + ) + + asyncio.create_task( + self._retry_after_delay( + job_id=str(job.id), + effective_priority=job.effective_priority, + scheduled_at=next_retry_at.timestamp(), + created_at=job.created_at.timestamp(), + delay=delay, + ) + ) + + else: + await db.execute( + update(Job) + .where(Job.id == job.id) + .values( + retry_count=new_retry_count, + worker_id=None, + updated_at=func.now(), + ) + ) + await db.flush() + await send_to_dlq(job.id, error_str[:1000], db) + + await _publish_sse_event( + { + "job_id": str(job.id), + "status": JobStatus.FAILED, + "is_dlq": True, + } + ) + + async def _retry_after_delay( + self, + job_id: str, + effective_priority: float, + scheduled_at: float, + created_at: float, + delay: float, + ) -> None: + """ + Background task: wait for the backoff delay then push the + job back onto the queue so it gets picked up again. + """ + await asyncio.sleep(delay) + try: + await self.queue.push( + job_id=job_id, + effective_priority=effective_priority, + scheduled_at=scheduled_at, + created_at=created_at, + ) + await _sync_job_to_redis(job_id, effective_priority) + logger.info("job_retry_queued", job_id=job_id, delay=delay) + except Exception as exc: + logger.error( + "job_retry_queue_error", + job_id=job_id, + error=str(exc), + ) + + async def _is_cancellation_requested( + self, + job_id: uuid.UUID, + db: AsyncSession, + ) -> bool: + """ + Check the cancellation_requested flag from the DB. + Called at checkpoints during processing. + """ + result = await db.execute( + select(Job.cancellation_requested).where(Job.id == job_id) + ) + return bool(result.scalar_one_or_none()) + + async def _mark_cancelled( + self, + job_id: uuid.UUID, + session: AsyncSession, + reason: str = "cancellation_requested", + ) -> None: + """Mark a job as cancelled and publish SSE event.""" + await session.execute( + update(Job) + .where(Job.id == job_id) + .values( + status=JobStatus.CANCELLED, + worker_id=None, + updated_at=func.now(), + ) + ) + + await self._log_event( + job_id=job_id, + event=LogEvent.JOB_CANCELLED, + message=(f"Job cancelled by worker {self.worker_id}. Reason: {reason}."), + metadata={"worker_id": self.worker_id, "reason": reason}, + db=session, + ) + await session.commit() + + logger.info( + "job_cancelled", + job_id=str(job_id), + worker_id=self.worker_id, + reason=reason, + ) + + await _publish_sse_event( + { + "job_id": str(job_id), + "status": JobStatus.CANCELLED, + "reason": reason, + } + ) + + async def _handle_recurrence( + self, + job: Job, + session: AsyncSession, + ) -> None: + """ + Schedule the next run of a recurring job. + """ + + if not job.interval_seconds: + return + + next_run = datetime.now(UTC) + timedelta(seconds=float(job.interval_seconds)) + + new_job = Job( + type=job.type, + payload=job.payload, + priority=job.priority, + effective_priority=float(job.priority), + status=JobStatus.PENDING, + scheduled_at=next_run, + interval_seconds=job.interval_seconds, + max_retries=job.max_retries, + retry_count=0, + ) + session.add(new_job) + await session.flush() + + await self._log_event( + job_id=new_job.id, + event=LogEvent.RECURRING_SCHEDULED, + message=( + f"Recurring job scheduled. Next run at {next_run.isoformat()}. " + f"Parent job: {job.id}." + ), + metadata={ + "parent_job_id": str(job.id), + "interval_seconds": job.interval_seconds, + "next_run": next_run.isoformat(), + }, + db=session, + ) + await session.commit() + + logger.info( + "recurring_job_scheduled", + parent_job_id=str(job.id), + new_job_id=str(new_job.id), + next_run=next_run.isoformat(), + interval_seconds=job.interval_seconds, + ) + + async def _load_job( + self, + job_id: str, + db: AsyncSession, + ) -> Job | None: + """Load a job by string ID. Returns None if not found.""" + try: + parsed_id = uuid.UUID(job_id) + except ValueError: + logger.error("invalid_job_id", job_id=job_id) + return None + + result = await db.execute( + select(Job).where( + Job.id == parsed_id, + Job.deleted_at.is_(None), + ) + ) + return result.scalar_one_or_none() + + async def _log_event( + self, + job_id: uuid.UUID, + event: str, + message: str, + metadata: dict, + db: AsyncSession, + ) -> None: + """Write a structured log entry to the job_logs table.""" + log_entry = JobLog( + job_id=job_id, + event=event, + message=message, + metadata_=metadata, + ) + db.add(log_entry) + await db.flush() + + @staticmethod + def _elapsed_ms(start: datetime) -> int: + """Return elapsed milliseconds since start.""" + delta = datetime.now(UTC) - start + return int(delta.total_seconds() * 1000) diff --git a/worker/worker.py b/worker/worker.py new file mode 100644 index 0000000..366ea4f --- /dev/null +++ b/worker/worker.py @@ -0,0 +1,302 @@ +import asyncio +import os +import signal +import sys +import uuid + +import redis.asyncio as aioredis + +from app.core.config import settings +from app.core.logger import get_logger, setup_logging +from app.queues.heapq import HeapQueue +from app.queues.timing_wheel import TimingWheel +from worker.processor import JobProcessor + +setup_logging() +logger = get_logger(__name__) + +WORKER_REGISTRY_KEY = "flint:workers:{worker_id}" +WORKER_CONTROL_CHANNEL = "flint:worker:control:{worker_id}" +QUEUE_KEY = "flint:queue" +HEARTBEAT_TTL = 60 # seconds — key expires if worker dies without deregistering +HEARTBEAT_INTERVAL = 30 # seconds between heartbeat refreshes + + +class FlintWorker: + def __init__(self) -> None: + self.worker_id = settings.WORKER_ID or f"worker-{str(uuid.uuid4())[:8]}" + self.running = False + self._redis: aioredis.Redis + self._queue: HeapQueue | TimingWheel | None = None + self._processor: JobProcessor + self._shutdown_event = asyncio.Event() + + async def start(self) -> None: + """ + Main entry point. Initialises resources and starts all loops. + """ + self.running = True + + self._redis = aioredis.from_url( + settings.REDIS_URL, + decode_responses=True, + encoding="utf-8", + ) + self._redis_pubsub = aioredis.from_url( + settings.REDIS_URL, + decode_responses=True, + encoding="utf-8", + socket_timeout=None, + socket_connect_timeout=5, + ) + + self._queue = await self._build_queue() + + await self._load_queue_from_redis() + + self._processor = JobProcessor( + worker_id=self.worker_id, + queue=self._queue, + ) + + await self._register() + + loop = asyncio.get_event_loop() + for sig in (signal.SIGTERM, signal.SIGINT): + loop.add_signal_handler(sig, self._handle_shutdown_signal) + + logger.info( + "worker_started", + worker_id=self.worker_id, + queue_strategy=self._queue.__class__.__name__, + ) + try: + await asyncio.gather( + self._poll_loop(), + self._heartbeat_loop(), + self._control_channel_loop(), + ) + finally: + await self._redis_pubsub.aclose() + await self._redis.aclose() + + async def _poll_loop(self) -> None: + """ + Core poll loop. Pops job IDs from the queue and processes them. + Sleeps for WORKER_POLL_INTERVAL when the queue is empty. + """ + while self.running: + try: + new_queue = await self._maybe_switch_queue() + if new_queue: + self._queue = new_queue + self._processor.queue = new_queue + await self._load_queue_from_redis() + + job_id = await self._pop_next() + if job_id: + await self._processor.process(job_id) + else: + await asyncio.sleep(settings.WORKER_POLL_INTERVAL) + + except asyncio.CancelledError: + break + except Exception as exc: + logger.error( + "worker_poll_error", + worker_id=self.worker_id, + error=str(exc), + ) + await asyncio.sleep(1.0) + + await self._deregister() + logger.info("worker_stopped", worker_id=self.worker_id) + + async def _pop_next(self) -> str | None: + """ + Pop the next job_id from the Redis sorted set. + Uses ZPOPMIN to atomically remove and return the lowest-score member. + """ + items = await self._redis.zpopmin(QUEUE_KEY, 1) + if not items: + return None + + job_id = items[0] # float(items[1]) if len(items) > 1 else float(items[0][1]), + + # Handle both tuple and flat list responses + if isinstance(items[0], (list, tuple)): + job_id = items[0][0] + # score = float(items[0][1]) + else: + job_id = items[0] + # score = float(items[1]) + + return str(job_id) + + async def _build_queue(self) -> HeapQueue | TimingWheel: + """Build the correct queue based on scheduler_strategy setting.""" + strategy = await self._get_strategy() + if strategy == "timing_wheel": + logger.info("queue_strategy", strategy="timing_wheel") + return TimingWheel() + logger.info("queue_strategy", strategy="heap") + return HeapQueue() + + async def _maybe_switch_queue(self) -> HeapQueue | TimingWheel | None: + """ + Check if the scheduler_strategy setting has changed. + If it has, return a new queue instance. Otherwise return None. + """ + current_strategy = ( + self._queue.__class__.__name__.lower() + .replace("heapqueue", "heap") + .replace("timingwheel", "timing_wheel") + ) + + new_strategy = await self._get_strategy() + + if new_strategy != current_strategy: + logger.info( + "queue_strategy_switched", + from_strategy=current_strategy, + to_strategy=new_strategy, + worker_id=self.worker_id, + ) + if new_strategy == "timing_wheel": + return TimingWheel() + return HeapQueue() + return None + + async def _get_strategy(self) -> str: + """Read scheduler_strategy from Redis cache or fall back to 'heap'.""" + try: + from app.db.session import AsyncSessionLocal + from app.services.settings import get_scheduler_strategy + + async with AsyncSessionLocal() as session: + return await get_scheduler_strategy(session) + except Exception: + return "heap" + + async def _load_queue_from_redis(self) -> None: + """ + On startup or strategy switch, load all pending job IDs from + the Redis sorted set into the local in-memory queue. + This populates the heap so the worker has the full queue state. + """ + import time + + try: + # ZRANGE with WITHSCORES returns [member, score, member, score, ...] + items = await self._redis.zrange(QUEUE_KEY, 0, -1, withscores=True) + count = 0 + for i in range(0, len(items), 2): + job_id = items[i] + score = float(items[i + 1]) + now = time.time() + await self._queue.push( + job_id=job_id, + effective_priority=score, + scheduled_at=now, + created_at=now, + ) + count += 1 + logger.info( + "queue_loaded_from_redis", + worker_id=self.worker_id, + job_count=count, + ) + except Exception as exc: + logger.error( + "queue_load_error", + worker_id=self.worker_id, + error=str(exc), + ) + + async def _heartbeat_loop(self) -> None: + """ + Refresh the worker's Redis registration key every HEARTBEAT_INTERVAL. + If the worker dies, the key expires after HEARTBEAT_TTL seconds + and the API reflects the worker as gone. + """ + while self.running: + try: + await self._register() + except Exception as exc: + logger.error( + "heartbeat_error", + worker_id=self.worker_id, + error=str(exc), + ) + await asyncio.sleep(HEARTBEAT_INTERVAL) + + async def _register(self) -> None: + """Write/refresh the worker registration key in Redis.""" + key = WORKER_REGISTRY_KEY.format(worker_id=self.worker_id) + await self._redis.set(key, "active", ex=HEARTBEAT_TTL) + + async def _deregister(self) -> None: + """Remove the worker registration key on clean shutdown.""" + key = WORKER_REGISTRY_KEY.format(worker_id=self.worker_id) + await self._redis.delete(key) + logger.info("worker_deregistered", worker_id=self.worker_id) + + async def _control_channel_loop(self) -> None: + """ + Subscribe to the worker's Redis control channel. + The API publishes 'stop' or 'restart' signals here. + This allows workers to be controlled from the dashboard. + """ + channel = WORKER_CONTROL_CHANNEL.format(worker_id=self.worker_id) + + while self.running: + pubsub = self._redis_pubsub.pubsub() + await pubsub.subscribe(channel) + try: + async for message in pubsub.listen(): + if not self.running: + return + if message["type"] != "message": + continue + command = message.get("data", "").strip().lower() + logger.info( + "worker_control_command", + worker_id=self.worker_id, + command=command, + ) + if command == "stop": + self._handle_shutdown_signal() + return + elif command == "restart": + logger.info("worker_restarting", worker_id=self.worker_id) + self._handle_shutdown_signal() + os.execv(sys.executable, [sys.executable] + sys.argv) + except asyncio.CancelledError: + return + except Exception as exc: + logger.warning( + "control_channel_reconnecting", + worker_id=self.worker_id, + error=str(exc), + ) + await asyncio.sleep(2) + finally: + await pubsub.unsubscribe(channel) + await pubsub.aclose() + + def _handle_shutdown_signal(self) -> None: + """ + Handle SIGTERM / SIGINT. + """ + if self.running: + logger.info( + "worker_shutdown_signal", + worker_id=self.worker_id, + ) + self.running = False + self._shutdown_event.set() + + +if __name__ == "__main__": + worker = FlintWorker() + asyncio.run(worker.start())