gitmind-backend / services /task_queue.py
Ak001z's picture
auto-deploy from CI (d604b09)
da0d5ba verified
Raw
History Blame Contribute Delete
9.34 kB
"""Background task queue β€” Redis List backend.
Enqueue: LPUSH to 'queue:tasks'.
Dequeue: BRPOP (blocking pop) from 'queue:tasks'.
State: Redis hash 'task:{id}' with TTL.
"""
import asyncio
import contextlib
import os
import time
import uuid
from collections.abc import Awaitable, Callable
import orjson
import structlog
from prometheus_client import Counter
from services.cache import get_redis
logger = structlog.get_logger(__name__)
_TASK_TTL = 60 * 10
_TASK_TIMEOUT = 300
_MAX_RETRIES = 3
_WORKER_COUNT = int(os.getenv("TASK_WORKERS", "3"))
_QUEUE_KEY = "queue:tasks"
_TASK_PREFIX = "task:"
_worker_tasks: list[asyncio.Task] = []
_prune_task: asyncio.Task | None = None
tasks_submitted = Counter("tasks_submitted_total", "Background tasks submitted")
tasks_completed = Counter("tasks_completed_total", "Background tasks completed", ["status"])
def _persist_task(task: dict) -> None:
"""Persist task state to Redis."""
r = get_redis()
if r:
r.set(
f"{_TASK_PREFIX}{task['id']}",
orjson.dumps(task, default=str).decode(),
ex=_TASK_TTL,
)
async def _worker():
"""Worker: BRPOP from queue, execute task, update state."""
r = get_redis()
if not r:
logger.error("No Redis β€” worker cannot start")
return
while True:
try:
result = await r.brpop(_QUEUE_KEY, timeout=5)
if result is None:
continue
_queue_key, payload = result
except asyncio.CancelledError:
raise
except Exception as exc:
logger.warning("BRPOP failed", error=str(exc)[:100])
await asyncio.sleep(1)
continue
try:
entry = orjson.loads(payload)
task_id = entry["task_id"]
retries_left = entry.get("retries", _MAX_RETRIES)
except Exception as exc:
logger.warning("Invalid queue entry", error=str(exc)[:100])
continue
# Load task state
raw = await r.get(f"{_TASK_PREFIX}{task_id}")
if not raw:
continue
task = orjson.loads(raw)
task["status"] = "running"
_persist_task(task)
# Deserialize coroutine factory
# We store the task as a pickled callable via orjson (not supported).
# Instead, we look up a registry of known task types.
coro_factory = _TASK_REGISTRY.get(task.get("type"))
if coro_factory is None:
task["status"] = "failed"
task["error"] = f"Unknown task type: {task.get('type')}"
_persist_task(task)
tasks_completed.labels(status="failed").inc()
continue
try:
coro = coro_factory(task.get("args", {}))
result = await asyncio.wait_for(coro, timeout=_TASK_TIMEOUT)
task["result"] = result
task["status"] = "completed"
tasks_completed.labels(status="completed").inc()
except asyncio.CancelledError:
task["status"] = "cancelled"
raise
except TimeoutError:
if retries_left > 0:
task["retries"] = retries_left - 1
task["status"] = "pending"
delay = 2 ** (_MAX_RETRIES - retries_left)
# Re-enqueue with delay
await asyncio.sleep(delay)
await r.lpush(
_QUEUE_KEY,
orjson.dumps(
{
"task_id": task_id,
"type": task.get("type"),
"args": task.get("args", {}),
"retries": retries_left - 1,
}
).decode(),
)
logger.warning("Task timed out, retrying", task_id=task_id, retries_left=retries_left - 1)
else:
task["error"] = "Task timed out after retries"
task["status"] = "failed"
tasks_completed.labels(status="failed").inc()
logger.error("Task permanently failed (timeout)", task_id=task_id)
except Exception as exc:
if retries_left > 0:
task["retries"] = retries_left - 1
task["status"] = "pending"
delay = 2 ** (_MAX_RETRIES - retries_left)
await asyncio.sleep(delay)
await r.lpush(
_QUEUE_KEY,
orjson.dumps(
{
"task_id": task_id,
"type": task.get("type"),
"args": task.get("args", {}),
"retries": retries_left - 1,
}
).decode(),
)
logger.warning("Task failed, retrying", task_id=task_id, retries_left=retries_left - 1)
else:
task["error"] = str(exc)[:500]
task["status"] = "failed"
tasks_completed.labels(status="failed").inc()
logger.error("Task permanently failed", task_id=task_id, error=str(exc)[:200])
finally:
task["completed_at"] = time.time()
_persist_task(task)
# ── Task registry ──────────────────────────────────────────────────────────────
# Maps task type string β†’ async callable that returns a coroutine.
# Add new task types here as they are needed.
# ── Built-in task factories ────────────────────────────────────────────────────
async def _run_analysis_task(args: dict) -> dict:
"""Run the full analysis pipeline for a repo."""
from routes.analyze import _run_analysis
return await _run_analysis(args["repo_url"], args.get("gh_token"), args["session_id"])
async def _run_pr_impact_task(args: dict) -> dict:
"""Compute PR impact analysis."""
from routes.pr import _compute_pr_impact
return await _compute_pr_impact(
repo_url=args["repo_url"],
token=args.get("token"),
head_ref=args["head_ref"],
base_session=args["base_session"],
pr_meta=args["pr_meta"],
pr_diff=args["pr_diff"],
)
_TASK_REGISTRY: dict[str, Callable[..., Awaitable[dict]]] = {
"run_analysis": _run_analysis_task,
"compute_pr_impact": _run_pr_impact_task,
}
def register_task_type(name: str, factory):
"""Register a task type. ``factory(args_dict)`` must return a coroutine."""
_TASK_REGISTRY[name] = factory
def submit(task_type: str, args: dict | None = None) -> str:
"""Submit a background task to Redis queue.
``task_type`` must be registered via ``register_task_type``.
``args`` is a JSON-serializable dict passed to the factory.
Returns task_id.
"""
r = get_redis()
if not r:
raise RuntimeError("Redis unavailable β€” cannot submit tasks")
task_id = str(uuid.uuid4())
task = {
"id": task_id,
"type": task_type,
"args": args or {},
"status": "pending",
"created_at": time.time(),
"completed_at": None,
"result": None,
"error": None,
"retries": _MAX_RETRIES,
}
_persist_task(task)
entry = {"task_id": task_id, "type": task_type, "args": args or {}, "retries": _MAX_RETRIES}
r.lpush(_QUEUE_KEY, orjson.dumps(entry).decode())
tasks_submitted.inc()
logger.debug("Task submitted", task_id=task_id, type=task_type)
return task_id
def get_task(task_id: str) -> dict | None:
"""Load task state from Redis."""
r = get_redis()
if not r:
return None
raw = r.get(f"{_TASK_PREFIX}{task_id}")
if raw:
return orjson.loads(raw)
return None
async def _periodic_prune():
"""Remove expired task hashes."""
while True:
await asyncio.sleep(300)
r = get_redis()
if not r:
continue
try:
cursor = 0
while True:
cursor, keys = await r.scan(cursor, match=f"{_TASK_PREFIX}*", count=100)
for key in keys:
ttl = await r.ttl(key)
if ttl == -1:
await r.expire(key, _TASK_TTL)
if cursor == 0:
break
except Exception as exc:
logger.warning("Prune scan failed", error=str(exc)[:100])
async def start():
"""Start worker tasks."""
global _worker_tasks, _prune_task
r = get_redis()
if not r:
logger.error("No Redis β€” task queue not started")
return
_worker_tasks = [asyncio.create_task(_worker()) for _ in range(_WORKER_COUNT)]
_prune_task = asyncio.create_task(_periodic_prune())
logger.info("Background task workers started", count=_WORKER_COUNT)
async def stop():
"""Cancel all workers."""
for t in _worker_tasks:
t.cancel()
if _prune_task:
_prune_task.cancel()
for t in _worker_tasks + [_prune_task]:
if t:
with contextlib.suppress(asyncio.CancelledError):
await t