mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 03:12:55 +00:00
fix(graphrag): give each extraction thread its own LLM
Provider-reported usage is kept on the LLM instance (_last_usage) and claimed by whichever call finishes next. With GRAPHRAG_EXTRACTION_WORKERS > 1 (the default is 8) every extraction thread shared one instance, so a call could claim another call's provider counts while its own fell back to the estimate: token_usage rows, and the cost they bill, could be attributed to the wrong call and summed wrong. Each pool thread now builds its own extraction LLM on first use. The calling thread's instance is still built up front, so a misconfigured model fails the run before any chunk is touched.
This commit is contained in:
1 parent
022cf69b49
commit
a20402a6dc
2 files changed
+75
-4
No files matched your search
@@ -227,6 +227,7 @@ def extract_graph_for_source(
|
||||
source's graph holds after the run — not how many upserts ran, which
|
||||
counts the same entity once per chunk it appears in.
|
||||
"""
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
from docsgpt.graphrag.store import GraphStore
|
||||
@@ -246,9 +247,24 @@ def extract_graph_for_source(
|
||||
|
||||
embedding = get_embeddings()
|
||||
|
||||
llm = _build_extraction_llm(
|
||||
_resolve_extraction_model(config), user, request_id
|
||||
)
|
||||
model_id = _resolve_extraction_model(config)
|
||||
# Built here first so a misconfigured model fails the run before any
|
||||
# chunk is touched; this instance serves the calling thread.
|
||||
thread_llm = threading.local()
|
||||
thread_llm.llm = _build_extraction_llm(model_id, user, request_id)
|
||||
|
||||
def _llm():
|
||||
"""This thread's extraction LLM.
|
||||
|
||||
Provider-reported usage is kept on the LLM instance (``_last_usage``)
|
||||
and claimed by whichever call finishes next, so two calls in flight on
|
||||
one instance can bill each other's tokens. Each pool thread therefore
|
||||
builds its own.
|
||||
"""
|
||||
llm = getattr(thread_llm, "llm", None)
|
||||
if llm is None:
|
||||
llm = thread_llm.llm = _build_extraction_llm(model_id, user, request_id)
|
||||
return llm
|
||||
|
||||
node_upserts = 0
|
||||
edges = 0
|
||||
@@ -284,7 +300,7 @@ def extract_graph_for_source(
|
||||
if not text:
|
||||
return chunk_id, "empty", None
|
||||
|
||||
extracted = _extract_chunk(llm, text, chunk_id)
|
||||
extracted = _extract_chunk(_llm(), text, chunk_id)
|
||||
if extracted is None:
|
||||
return chunk_id, "failed", None
|
||||
try:
|
||||
|
||||
@@ -606,6 +606,61 @@ class TestExtractionTokenUsage:
|
||||
assert built._request_id == "req-99"
|
||||
assert captured["model_id"] == "stub-model"
|
||||
|
||||
def test_concurrent_extraction_calls_never_share_an_llm(self, monkeypatch, stub_embedding):
|
||||
"""Provider usage is recorded on the LLM instance (``_last_usage``) and
|
||||
claimed by whichever call finishes next, so two calls in flight on one
|
||||
instance can bill each other's tokens. Each extraction thread needs its
|
||||
own instance."""
|
||||
import threading
|
||||
import time
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from docsgpt.core.settings import settings
|
||||
|
||||
payload = _extraction_json(
|
||||
entities=[{"name": "Ada", "type": "person", "description": "d"}],
|
||||
relationships=[],
|
||||
)
|
||||
|
||||
class _ThreadRecordingLLM:
|
||||
model_id = "stub-model"
|
||||
|
||||
def __init__(self):
|
||||
self.threads = set()
|
||||
|
||||
def gen(self, model=None, messages=None, **kwargs):
|
||||
self.threads.add(threading.get_ident())
|
||||
time.sleep(0.01) # keep calls overlapping
|
||||
return payload
|
||||
|
||||
built = []
|
||||
|
||||
def _create(*args, **kwargs):
|
||||
llm = _ThreadRecordingLLM()
|
||||
built.append(llm)
|
||||
return llm
|
||||
|
||||
monkeypatch.setattr(extraction_module.LLMCreator, "create_llm", staticmethod(_create))
|
||||
monkeypatch.setattr(settings, "GRAPHRAG_EXTRACTION_WORKERS", 4)
|
||||
store = MagicMock(name="GraphStore")
|
||||
store.pending_chunks.return_value = [f"c{i}" for i in range(8)]
|
||||
store.apply_chunk.return_value = (1, 0)
|
||||
store.count_nodes.return_value = 1
|
||||
monkeypatch.setattr("docsgpt.graphrag.store.GraphStore", lambda *a, **k: store)
|
||||
|
||||
summary = extract_graph_for_source(
|
||||
str(uuid.uuid4()),
|
||||
user="owner-1",
|
||||
chunks=[_chunk(f"c{i}", f"Ada, take {i}.") for i in range(8)],
|
||||
config=SourceConfig(),
|
||||
request_id="req-threads",
|
||||
)
|
||||
|
||||
assert summary["chunks_processed"] == 8
|
||||
used = [llm for llm in built if llm.threads]
|
||||
assert len(used) > 1, "calls did not run concurrently"
|
||||
assert all(len(llm.threads) == 1 for llm in used)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestModelResolution:
|
||||
|
||||
Reference in new issue
Block a user