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:
Alex committed 2026-09-19 14:33:14 +01:00
1 parent 022cf69b49
commit a20402a6dc
2 files changed
+75 -4

No files matched your search

+20 -4
View File
@@ -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:
+55
View File
@@ -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: