mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 03:12:55 +00:00
fix(worker): know you are in a worker from any thread, not only the task's
Celery records the executing task on the thread that runs it, so a thread that
task starts sees none. The embeddings client and read_document both decided
"am I in a worker?" from that alone, and from any other thread took the
web-process branch: dispatch to the worker they were running in and block on
the result. Celery refuses that get() ("Never call result.get() within a
task!"), so the embed failed and latched the 30s dispatch cooldown for every
caller after it; with joins allowed, read_document would instead wait on a
parsing queue only its own busy process serves.
Threads inside tasks are not hypothetical: per-source retrieval fans out to a
pool, so a scheduled or webhook agent searching several sources embedded from
pool threads. Graph extraction did too, which failed every chunk of a build.
in_worker() in celery_init answers for the whole process. Celery's
task_join_will_block is process-wide and set for every blocking pool (prefork,
solo, threads) -- exactly the condition under which dispatch-and-wait goes
wrong; eventlet/gevent leave it unset, so the task's own thread still counts
through current_worker_task. Verified with real workers on each blocking pool:
from a thread a task started, the old check dispatched and hit the error, the
new one embedded locally.
This commit is contained in:
1 parent
5e67f8c927
commit
15bdda8554
6 files changed
+117
-16
No files matched your search
@@ -17,8 +17,6 @@ import signal
|
||||
import threading
|
||||
from typing import Any, Callable, Dict, List, Optional
|
||||
|
||||
from celery import current_task
|
||||
|
||||
from docsgpt.agents.tools.artifact_ref import resolve_artifact_id
|
||||
from docsgpt.agents.tools.attachment_bridge import (
|
||||
AttachmentBridgeError,
|
||||
@@ -26,6 +24,7 @@ from docsgpt.agents.tools.attachment_bridge import (
|
||||
match_attachment,
|
||||
)
|
||||
from docsgpt.agents.tools.base import Tool
|
||||
from docsgpt.celery_init import in_worker
|
||||
from docsgpt.core.json_schema_utils import (
|
||||
JsonSchemaValidationError,
|
||||
normalize_json_schema_payload,
|
||||
@@ -229,9 +228,9 @@ class ReadDocumentTool(Tool):
|
||||
# (floored at DOCUMENT_PARSE_TIMEOUT).
|
||||
timeout = parse_timeout_for_size(self._input_size)
|
||||
|
||||
# ``current_task`` is a Celery proxy: truthy only while this runs inside a worker task,
|
||||
# falsy in the web process (the bare proxy is NOT identity-None, so test truthiness).
|
||||
if current_task:
|
||||
# Process-wide, not the thread-local ``current_task``: a thread a task starts has no
|
||||
# task of its own, and dispatching from there is the self-deadlock described above.
|
||||
if in_worker():
|
||||
from docsgpt.worker import run_parse_document
|
||||
|
||||
try:
|
||||
|
||||
@@ -172,6 +172,29 @@ def _run_version_check(*args, **kwargs):
|
||||
celery = make_celery()
|
||||
celery.config_from_object("docsgpt.celeryconfig")
|
||||
|
||||
|
||||
def in_worker() -> bool:
|
||||
"""True anywhere in a Celery worker process, on any thread.
|
||||
|
||||
``current_worker_task`` alone is not enough: Celery records the executing
|
||||
task on the thread that runs it, so a thread 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 only this busy process serves.
|
||||
|
||||
``task_join_will_block`` is process-wide and set for every blocking pool
|
||||
(prefork, solo, threads) — exactly the condition under which dispatching
|
||||
and waiting goes wrong. eventlet/gevent leave it unset, so the task's own
|
||||
thread still counts through ``current_worker_task``.
|
||||
|
||||
Returns:
|
||||
bool: Whether this call is running inside a worker process.
|
||||
"""
|
||||
from celery.result import task_join_will_block
|
||||
|
||||
return task_join_will_block() or celery.current_worker_task is not None
|
||||
|
||||
#: Task-name prefix the package carried before the rename to ``docsgpt``.
|
||||
|
||||
|
||||
|
||||
@@ -10,8 +10,9 @@ Celery and the vector comes back. The API pays a broker round trip per query
|
||||
and no resident model.
|
||||
|
||||
Inside a worker there is nothing to delegate to -- dispatching would queue work
|
||||
behind the task already running and wait on itself -- so a call made while a
|
||||
task is executing runs locally, on a model this process loads once and caches.
|
||||
behind the task already running and wait on itself -- so a call made anywhere in
|
||||
a worker process, including from a thread a task started, runs locally, on a
|
||||
model this process loads once and caches.
|
||||
``DOCUMENT_PARSE_QUEUE`` exists for the same reason on the parsing side.
|
||||
|
||||
Production deployments should point ``EMBEDDINGS_BASE_URL`` at a real embedding
|
||||
@@ -79,11 +80,11 @@ def _forget(result) -> None:
|
||||
|
||||
|
||||
def _in_worker() -> bool:
|
||||
"""True when a Celery task is executing in this process."""
|
||||
"""True anywhere in a Celery worker process -- on any thread, not only the task's."""
|
||||
try:
|
||||
from docsgpt.celery_init import celery
|
||||
from docsgpt.celery_init import in_worker
|
||||
|
||||
return celery.current_worker_task is not None
|
||||
return in_worker()
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
@@ -357,9 +357,9 @@ def test_malformed_json_schema_rejected_before_enqueue(monkeypatch):
|
||||
@pytest.mark.unit
|
||||
def test_dispatch_inline_when_in_worker(monkeypatch):
|
||||
_stub_repo(monkeypatch, found=True, conv="conv-1", run=None)
|
||||
# Inside a worker current_task is truthy -> parse inline, never enqueue (else the
|
||||
# Inside a worker -> parse inline, never enqueue (else the
|
||||
# parsing queue self-deadlocks the worker that also serves it).
|
||||
monkeypatch.setattr(rd, "current_task", object())
|
||||
monkeypatch.setattr(rd, "in_worker", lambda: True)
|
||||
|
||||
import docsgpt.api.user.tasks as tasks
|
||||
monkeypatch.setattr(
|
||||
@@ -387,8 +387,8 @@ def test_dispatch_inline_when_in_worker(monkeypatch):
|
||||
@pytest.mark.unit
|
||||
def test_dispatch_enqueues_when_not_in_worker(monkeypatch):
|
||||
_stub_repo(monkeypatch, found=True, conv="conv-1", run=None)
|
||||
# Web process: current_task falsy -> dispatch to the parsing queue, never inline.
|
||||
monkeypatch.setattr(rd, "current_task", None)
|
||||
# Web process -> dispatch to the parsing queue, never inline.
|
||||
monkeypatch.setattr(rd, "in_worker", lambda: False)
|
||||
captured = _patch_task(monkeypatch, payload={"status": "ok", "content": "queued", "truncated": False})
|
||||
|
||||
import docsgpt.worker as worker
|
||||
@@ -414,9 +414,9 @@ _TIMED_OUT = "document parsing timed out after"
|
||||
|
||||
|
||||
def _inline(monkeypatch, run_parse, *, timeout=0.2) -> ReadDocumentTool:
|
||||
"""Drive the inline branch (current_task truthy) with a patched parse window."""
|
||||
"""Drive the inline (in-worker) branch with a patched parse window."""
|
||||
_stub_repo(monkeypatch, found=True, conv="conv-1", run=None)
|
||||
monkeypatch.setattr(rd, "current_task", object())
|
||||
monkeypatch.setattr(rd, "in_worker", lambda: True)
|
||||
|
||||
import docsgpt.api.user.tasks as tasks
|
||||
monkeypatch.setattr(
|
||||
|
||||
@@ -276,3 +276,55 @@ class TestReclaimIsSkippedForEmbeds:
|
||||
from docsgpt.vectorstore.embeddings_delegated import EMBED_TASK
|
||||
|
||||
assert EMBED_TASK in _NO_RECLAIM_TASKS
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestInWorker:
|
||||
"""Whether code runs inside a worker must not depend on which thread asks.
|
||||
|
||||
Celery records the executing task on the thread that runs it, so a thread
|
||||
that task starts sees no task at all. Code deciding "am I in the worker?"
|
||||
from that alone takes the web-process branch there: it dispatches to the
|
||||
worker it is running in and blocks on the result, which Celery refuses
|
||||
("Never call result.get() within a task!") or, where joins are allowed,
|
||||
waits on a queue only this busy process serves.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _ask_from_a_new_thread():
|
||||
import threading
|
||||
|
||||
from docsgpt.celery_init import in_worker
|
||||
|
||||
seen = []
|
||||
thread = threading.Thread(target=lambda: seen.append(in_worker()))
|
||||
thread.start()
|
||||
thread.join()
|
||||
return seen[0]
|
||||
|
||||
def test_false_outside_a_worker(self):
|
||||
from docsgpt.celery_init import in_worker
|
||||
|
||||
assert in_worker() is False
|
||||
assert self._ask_from_a_new_thread() is False
|
||||
|
||||
def test_true_on_a_thread_started_inside_a_worker(self):
|
||||
# Blocking pools (prefork, solo, threads) mark the whole process as one
|
||||
# where joining a task would block; ``denied_join_result`` sets exactly
|
||||
# that flag.
|
||||
from celery.result import denied_join_result
|
||||
|
||||
with denied_join_result():
|
||||
assert self._ask_from_a_new_thread() is True
|
||||
|
||||
def test_true_on_the_task_thread_of_a_non_blocking_pool(self):
|
||||
# eventlet/gevent pools leave the process flag unset; the thread
|
||||
# running the task still knows it is in one.
|
||||
from unittest.mock import PropertyMock
|
||||
|
||||
from docsgpt.celery_init import celery, in_worker
|
||||
|
||||
with patch.object(
|
||||
type(celery), "current_worker_task", new_callable=PropertyMock, return_value=object()
|
||||
):
|
||||
assert in_worker() is True
|
||||
@@ -72,6 +72,32 @@ class TestInsideAWorker:
|
||||
assert vector == [1.0, 2.0]
|
||||
celery.send_task.assert_not_called()
|
||||
|
||||
def test_a_thread_started_inside_the_worker_embeds_locally(self):
|
||||
"""The task's own thread is not the only one in a worker.
|
||||
|
||||
Graph extraction and per-source retrieval both fan out to thread pools
|
||||
inside tasks. The check used to read the task off the current thread
|
||||
only, so from those threads it dispatched to the worker it was running
|
||||
in -- and Celery refuses that ``get()`` inside a worker, failing the
|
||||
call and latching the 30s dispatch cooldown for every caller after it.
|
||||
"""
|
||||
from celery.result import denied_join_result
|
||||
|
||||
from docsgpt.celery_init import celery
|
||||
|
||||
local = MagicMock()
|
||||
local.embed_documents.return_value = [[1.0, 2.0]]
|
||||
client = DelegatedEmbeddings("some/model")
|
||||
vectors = []
|
||||
with denied_join_result():
|
||||
with patch("docsgpt.vectorstore.base.build_local_embeddings", return_value=local):
|
||||
with patch.object(celery, "send_task") as send_task:
|
||||
thread = threading.Thread(target=lambda: vectors.append(client.embed_query("hi")))
|
||||
thread.start()
|
||||
thread.join()
|
||||
assert vectors == [[1.0, 2.0]]
|
||||
send_task.assert_not_called()
|
||||
|
||||
def test_the_local_model_is_built_once(self):
|
||||
local = MagicMock()
|
||||
local.embed_documents.return_value = [[1.0]]
|
||||
|
||||
Reference in new issue
Block a user