Files
DocsGPT/docsgpt/celery_init.py
T
Alex 877609dfb6 fix(worker): record worker state at startup, for every pool
in_worker() read task_join_will_block, which eventlet and gevent leave unset,
and the task's current_worker_task, which they scope to one greenlet -- so a
greenlet a task spawned in those pools still took the web-process branch and
dispatched to its own worker, where it could queue behind its parent and time
out.

The worker's own startup now records it: worker_init fires in every worker's
main process before the pool starts (where solo, threads, eventlet and gevent
run tasks, and what prefork children fork from), and worker_process_init in
each prefork child. worker_ready would be too late -- prefork children are
forked before it fires. The two existing checks stay for anything that runs
tasks without that startup.
2026-09-19 14:33:14 +01:00

259 lines
9.0 KiB
Python

import ctypes
import gc
import inspect
import logging
import sys
import threading
from celery import Celery
from docsgpt.core import log_context
from docsgpt.core.settings import settings
from celery.signals import (
celeryd_after_setup,
setup_logging,
task_postrun,
task_prerun,
worker_init,
worker_process_init,
worker_ready,
)
def make_celery(app_name=__name__):
celery = Celery(
app_name,
broker=settings.CELERY_BROKER_URL,
backend=settings.CELERY_RESULT_BACKEND,
)
celery.conf.update(settings)
return celery
@setup_logging.connect
def config_loggers(*args, **kwargs):
from docsgpt.core.logging_config import setup_logging
setup_logging()
@worker_process_init.connect
def _dispose_db_engine_on_fork(*args, **kwargs):
"""Dispose the SQLAlchemy engine pool in each forked Celery worker.
SQLAlchemy connection pools are not fork-safe: file descriptors shared
between the parent and a forked worker will corrupt the pool. Disposing
on ``worker_process_init`` gives every worker its own fresh pool on
first use.
Imported lazily so Celery workers that don't touch Postgres (or where
``POSTGRES_URI`` is unset) don't fail at startup.
"""
try:
from docsgpt.storage.db.engine import dispose_engine
except Exception:
return
dispose_engine()
# Most tasks in this repo accept ``user`` where the log context wants
# ``user_id``; map task parameter names to context keys explicitly.
_TASK_PARAM_TO_CTX_KEY: dict[str, str] = {
"user": "user_id",
"user_id": "user_id",
"agent_id": "agent_id",
"conversation_id": "conversation_id",
}
_task_log_tokens: dict[str, object] = {}
@task_prerun.connect
def _bind_task_log_context(task_id, task, args, kwargs, **_):
# Resolve task args by parameter name — nearly every task in this repo
# is called positionally, so ``kwargs.get('user')`` would bind nothing.
ctx = {"activity_id": task_id}
try:
sig = inspect.signature(task.run)
bound = sig.bind_partial(*args, **kwargs).arguments
except (TypeError, ValueError):
bound = dict(kwargs)
for param_name, value in bound.items():
ctx_key = _TASK_PARAM_TO_CTX_KEY.get(param_name)
if ctx_key and value:
ctx[ctx_key] = value
_task_log_tokens[task_id] = log_context.bind(**ctx)
@task_postrun.connect
def _unbind_task_log_context(task_id, **_):
# ``task_postrun`` fires on both success and failure. Required for
# Celery: unlike the Flask path, tasks aren't isolated in their own
# ``copy_context().run(...)``, so a missing reset would leak the
# bind onto the next task on the same worker.
token = _task_log_tokens.pop(task_id, None)
if token is None:
return
try:
log_context.reset(token)
except ValueError:
# task_prerun and task_postrun ran on different threads (non-default
# Celery pool); the token isn't valid in this context. Drop it.
logging.getLogger(__name__).debug(
"log_context reset skipped for task %s", task_id
)
def _trim_native_heap() -> None:
"""Return freed glibc heap pages to the OS (Linux only; no-op elsewhere)."""
# docling/torch parsing makes large transient allocations; glibc keeps the
# freed pages in per-thread malloc arenas rather than returning them, so a
# long-lived worker child's RSS only ever climbs. malloc_trim hands them
# back. The symbol is glibc-only — absent in macOS libc.
if not sys.platform.startswith("linux"):
return
try:
ctypes.CDLL("libc.so.6").malloc_trim(0)
except (OSError, AttributeError):
pass
# Tasks that allocate almost nothing and run on a latency-sensitive path, so
# the reclaim below costs far more than it recovers. Query embedding is one:
# measured at ~86 ms for the collect against ~8 ms for the embed itself on a
# worker holding the ONNX model, i.e. a 9x slowdown of the whole round trip.
_NO_RECLAIM_TASKS = frozenset({"docsgpt.vectorstore.embeddings_tasks.embed_texts"})
LEGACY_TASK_PREFIX = "application."
@task_postrun.connect
def _reclaim_memory_after_task(task=None, **kwargs):
"""Drop per-task allocations so the prefork child's RSS doesn't ratchet.
Skipped for the tasks in :data:`_NO_RECLAIM_TASKS`. This exists for the
large transient allocations docling/torch parsing makes; running a full
generational collect after a task that allocated a few kilobytes just
charges the next task for walking the whole heap.
"""
name = getattr(task, "name", None)
if isinstance(name, str) and name.startswith(LEGACY_TASK_PREFIX):
name = "docsgpt." + name[len(LEGACY_TASK_PREFIX):]
if name in _NO_RECLAIM_TASKS:
return
gc.collect()
torch = sys.modules.get("torch")
if torch is not None:
try:
if torch.cuda.is_available():
torch.cuda.empty_cache()
except Exception:
pass
_trim_native_heap()
@worker_ready.connect
def _run_version_check(*args, **kwargs):
"""Kick off the anonymous version check on worker startup.
Runs in a daemon thread so a slow endpoint or bad DNS never holds
up the worker becoming ready for tasks. The check itself is
fail-silent (see ``docsgpt.updates.version_check.run_check``);
this handler's only job is to launch it and get out of the way.
Import is lazy so the symbol resolution never fires at module
import time — consistent with the ``_dispose_db_engine_on_fork``
pattern above.
"""
try:
from docsgpt.updates.version_check import run_check
except Exception:
return
threading.Thread(target=run_check, name="version-check", daemon=True).start()
celery = make_celery()
celery.config_from_object("docsgpt.celeryconfig")
#: Set once this process starts as a worker; see :func:`_mark_worker_process`.
_IS_WORKER_PROCESS = False
@worker_init.connect
@worker_process_init.connect
def _mark_worker_process(*args, **kwargs):
"""Record that this process runs tasks, for :func:`in_worker`.
``worker_init`` fires in every worker's main process before its pool
starts: that is where solo, threads, eventlet and gevent run tasks, and
what prefork children fork from. ``worker_process_init`` covers prefork
children however they were started.
"""
global _IS_WORKER_PROCESS
_IS_WORKER_PROCESS = True
def in_worker() -> bool:
"""True anywhere in a Celery worker process, on any thread or greenlet.
``current_worker_task`` alone is not enough: Celery records the executing
task on the thread (or greenlet) that runs it, so one the task starts sees
none and would take the web-process branch — dispatching to the worker it
is running in and blocking on the result. Celery refuses that ``get()``
("Never call result.get() within a task!"), or, where joins are allowed,
it waits on a queue that only this busy process may be able to serve.
The worker's own startup (:func:`_mark_worker_process`) answers for every
pool. ``task_join_will_block`` — process-wide, set for every blocking pool
— and the task's own ``current_worker_task`` still count for a process
that runs tasks without having gone through that startup.
Returns:
bool: Whether this call is running inside a worker process.
"""
from celery.result import task_join_will_block
return (
_IS_WORKER_PROCESS
or task_join_will_block()
or celery.current_worker_task is not None
)
#: Task-name prefix the package carried before the rename to ``docsgpt``.
def register_legacy_task_names(app: Celery) -> int:
"""Make every ``docsgpt.*`` task answer to its old ``application.*`` name too.
Messages queued by the previous release carry the old names; without the
alias a worker on this release rejects them as unregistered. Each alias is
a distinct task object (a subclass carrying the old name), not the same
object under a second key: Celery builds its execution tracer per task
object, and one object under two names would log every run under
whichever name was traced last. Kept for one release, together with the
``application`` import alias.
Returns:
The number of aliases added.
"""
added = 0
for name, task in list(app.tasks.items()):
if not name.startswith("docsgpt."):
continue
legacy = LEGACY_TASK_PREFIX + name[len("docsgpt."):]
if legacy in app.tasks:
continue
base = type(task)
legacy_cls = type(base.__name__, (base,), {"name": legacy, "__module__": base.__module__, "__doc__": base.__doc__})
app.register_task(legacy_cls())
added += 1
return added
@celeryd_after_setup.connect
def _alias_legacy_task_names(sender=None, instance=None, **kwargs):
"""Register the pre-rename task names once the worker has loaded its tasks."""
app = getattr(instance, "app", None) or celery
added = register_legacy_task_names(app)
if added:
logging.getLogger(__name__).info("Registered %d legacy 'application.*' task-name aliases", added)