Skip to content

Commit c7ea8ab

Browse files
Sebastian BraunCopilot
andcommitted
fix(agent): stream LLM completions to avoid gateway idle-timeout
Corporate LLM gateways (e.g. AI.proxy on AWS) enforce an idle timeout on buffered (non-streaming) requests, so a long-running compile step can hit a Gateway Timeout even though the provider would have eventually finished. Switch _llm_call() and _llm_call_async() in openkb/agent/compiler.py to litellm.completion()/acompletion() with stream=True: streaming keeps bytes flowing over the connection, so idle-timeout gateways never see a silent connection. Chunks are merged back into the existing response shape via a new _merge_stream_chunks() helper, using LiteLLM's own litellm.stream_chunk_builder() for genuine multi-chunk streams. An exception raised mid-stream propagates as a complete failure (list() never returns a partial buffer), matching prior all-or-nothing behavior. Adapts the compiler test mocks (_mock_completion/_mock_acompletion and a handful of inline mocks) to return a single-chunk fake stream, plus the litellm.completion/acompletion mocks in test_llm_timeout.py. Resolves #235. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
1 parent ff54396 commit c7ea8ab

3 files changed

Lines changed: 91 additions & 83 deletions

File tree

‎openkb/agent/compiler.py‎

Lines changed: 45 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -397,6 +397,23 @@ class TruncatedResponseError(Exception):
397397
treat truncation as a failure (so a partial page is skipped, not written)."""
398398

399399

400+
def _merge_stream_chunks(chunks: list, messages: list[dict]):
401+
"""Merge streamed LLM chunks back into a single, non-streaming response.
402+
403+
Genuine LiteLLM stream chunks only ever carry a ``.delta`` (never a
404+
``.message``), so a real multi-chunk stream is merged via LiteLLM's own
405+
:func:`litellm.stream_chunk_builder`. A single chunk that already looks
406+
like a complete, non-streaming ``ModelResponse`` (exposing ``.message``)
407+
is used as-is — there's nothing left to merge, and it lets test doubles
408+
fake a one-shot response without simulating LiteLLM's internal delta
409+
format.
410+
"""
411+
choices = getattr(chunks[0], "choices", None) or []
412+
if len(chunks) == 1 and choices and hasattr(choices[0], "message"):
413+
return chunks[0]
414+
return litellm.stream_chunk_builder(chunks, messages=messages)
415+
416+
400417
def _llm_call(
401418
model: str,
402419
messages: list[dict],
@@ -406,7 +423,15 @@ def _llm_call(
406423
bundle=None,
407424
**kwargs,
408425
) -> str:
409-
"""Single LLM call with animated progress and debug logging."""
426+
"""Single LLM call with animated progress and debug logging.
427+
428+
Uses ``stream=True``: some corporate LLM gateways enforce an idle
429+
timeout on buffered (non-streaming) requests, which a long-running
430+
completion can hit before the response is ever sent. Streaming keeps
431+
bytes flowing over the connection so that timeout never fires; the
432+
chunks are merged back into a single response via
433+
:func:`_merge_stream_chunks` so callers see the same shape as before.
434+
"""
410435
messages = _prepare_messages(model, messages)
411436
extra_headers = bundle.extra_headers if bundle is not None else get_extra_headers()
412437
if extra_headers:
@@ -417,6 +442,7 @@ def _llm_call(
417442
if bundle is not None:
418443
kwargs.setdefault("api_key", bundle.api_key)
419444
kwargs.setdefault("base_url", bundle.base_url)
445+
kwargs.setdefault("stream_options", {"include_usage": True})
420446
logger.debug("LLM request [%s]:\n%s", step_name, _fmt_messages(messages))
421447
if kwargs:
422448
logger.debug("LLM kwargs [%s]: %s", step_name, kwargs)
@@ -425,7 +451,11 @@ def _llm_call(
425451
spinner.start()
426452
t0 = time.time()
427453

428-
response = litellm.completion(model=model, messages=messages, **kwargs)
454+
stream = litellm.completion(model=model, messages=messages, stream=True, **kwargs)
455+
chunks = list(stream)
456+
if not chunks:
457+
raise RuntimeError(f"LLM [{step_name}] stream produced no chunks")
458+
response = _merge_stream_chunks(chunks, messages)
429459
content = response.choices[0].message.content or ""
430460
truncated = _warn_if_truncated(response, step_name, kwargs.get("max_tokens"))
431461

@@ -449,7 +479,10 @@ async def _llm_call_async(
449479
bundle=None,
450480
**kwargs,
451481
) -> str:
452-
"""Async LLM call with timing output and debug logging."""
482+
"""Async LLM call with timing output and debug logging.
483+
484+
See ``_llm_call`` for why ``stream=True`` is used.
485+
"""
453486
messages = _prepare_messages(model, messages)
454487
extra_headers = bundle.extra_headers if bundle is not None else get_extra_headers()
455488
if extra_headers:
@@ -460,13 +493,21 @@ async def _llm_call_async(
460493
if bundle is not None:
461494
kwargs.setdefault("api_key", bundle.api_key)
462495
kwargs.setdefault("base_url", bundle.base_url)
496+
kwargs.setdefault("stream_options", {"include_usage": True})
463497
logger.debug("LLM request [%s]:\n%s", step_name, _fmt_messages(messages))
464498
if kwargs:
465499
logger.debug("LLM kwargs [%s]: %s", step_name, kwargs)
466500

467501
t0 = time.time()
468502

469-
response = await litellm.acompletion(model=model, messages=messages, **kwargs)
503+
stream = await litellm.acompletion(model=model, messages=messages, stream=True, **kwargs)
504+
if hasattr(stream, "__aiter__"):
505+
chunks = [chunk async for chunk in stream]
506+
else:
507+
chunks = list(stream)
508+
if not chunks:
509+
raise RuntimeError(f"LLM [{step_name}] stream produced no chunks")
510+
response = _merge_stream_chunks(chunks, messages)
470511
content = response.choices[0].message.content or ""
471512
truncated = _warn_if_truncated(response, step_name, kwargs.get("max_tokens"))
472513

‎tests/test_compiler.py‎

Lines changed: 33 additions & 73 deletions
Original file line numberDiff line numberDiff line change
@@ -1112,36 +1112,45 @@ def test_frontmatter_without_sources_line_gets_one_inserted(self, tmp_path):
11121112
assert "[[summaries/new-doc]]" in text
11131113

11141114

1115+
def _mock_response(content, finish_reason: str = "stop") -> MagicMock:
1116+
"""Build a fake, already-complete LLM response (single-chunk stream).
1117+
1118+
``_llm_call``/``_llm_call_async`` now call ``litellm.completion``/
1119+
``acompletion`` with ``stream=True`` and merge the resulting chunks back
1120+
into one response (see ``_merge_stream_chunks``). Exposing ``.message``
1121+
(rather than the ``.delta`` a genuine stream chunk carries) tells
1122+
``_merge_stream_chunks`` this single chunk *is* the final response, so it
1123+
is used as-is without needing to fake LiteLLM's internal delta format.
1124+
"""
1125+
mock_resp = MagicMock()
1126+
mock_resp.choices = [MagicMock()]
1127+
mock_resp.choices[0].message.content = content
1128+
mock_resp.choices[0].finish_reason = finish_reason
1129+
mock_resp.usage = MagicMock(prompt_tokens=100, completion_tokens=50)
1130+
mock_resp.usage.prompt_tokens_details = None
1131+
return mock_resp
1132+
1133+
11151134
def _mock_completion(responses: list[str]):
1116-
"""Create a mock for litellm.completion that returns responses in order."""
1135+
"""Create a mock for litellm.completion returning a single-chunk stream."""
11171136
call_count = {"n": 0}
11181137

11191138
def side_effect(*args, **kwargs):
11201139
idx = min(call_count["n"], len(responses) - 1)
11211140
call_count["n"] += 1
1122-
mock_resp = MagicMock()
1123-
mock_resp.choices = [MagicMock()]
1124-
mock_resp.choices[0].message.content = responses[idx]
1125-
mock_resp.usage = MagicMock(prompt_tokens=100, completion_tokens=50)
1126-
mock_resp.usage.prompt_tokens_details = None
1127-
return mock_resp
1141+
return [_mock_response(responses[idx])]
11281142

11291143
return side_effect
11301144

11311145

11321146
def _mock_acompletion(responses: list[str]):
1133-
"""Create an async mock for litellm.acompletion."""
1147+
"""Create an async mock for litellm.acompletion returning a single-chunk stream."""
11341148
call_count = {"n": 0}
11351149

11361150
async def side_effect(*args, **kwargs):
11371151
idx = min(call_count["n"], len(responses) - 1)
11381152
call_count["n"] += 1
1139-
mock_resp = MagicMock()
1140-
mock_resp.choices = [MagicMock()]
1141-
mock_resp.choices[0].message.content = responses[idx]
1142-
mock_resp.usage = MagicMock(prompt_tokens=100, completion_tokens=50)
1143-
mock_resp.usage.prompt_tokens_details = None
1144-
return mock_resp
1153+
return [_mock_response(responses[idx])]
11451154

11461155
return side_effect
11471156

@@ -1342,15 +1351,7 @@ def sync_side_effect(*args, **kwargs):
13421351
sync_call_count["n"] += 1
13431352
if idx == 2: # the summary-rewrite call
13441353
raise RuntimeError("simulated API failure")
1345-
mock_resp = MagicMock()
1346-
mock_resp.choices = [MagicMock()]
1347-
mock_resp.choices[0].message.content = [
1348-
summary_response,
1349-
plan_response,
1350-
][idx]
1351-
mock_resp.usage = MagicMock(prompt_tokens=1, completion_tokens=1)
1352-
mock_resp.usage.prompt_tokens_details = None
1353-
return mock_resp
1354+
return [_mock_response([summary_response, plan_response][idx])]
13541355

13551356
with patch("openkb.agent.compiler.litellm") as mock_litellm:
13561357
mock_litellm.completion = MagicMock(side_effect=sync_side_effect)
@@ -1507,21 +1508,11 @@ async def test_short_doc_marks_doc_and_summary(self, tmp_path):
15071508
def sync_side_effect(*args, **kwargs):
15081509
captured_sync_calls.append(kwargs["messages"])
15091510
idx = min(len(captured_sync_calls) - 1, len(sync_responses) - 1)
1510-
mock_resp = MagicMock()
1511-
mock_resp.choices = [MagicMock()]
1512-
mock_resp.choices[0].message.content = sync_responses[idx]
1513-
mock_resp.usage = MagicMock(prompt_tokens=1, completion_tokens=1)
1514-
mock_resp.usage.prompt_tokens_details = None
1515-
return mock_resp
1511+
return [_mock_response(sync_responses[idx])]
15161512

15171513
async def async_side_effect(*args, **kwargs):
15181514
captured_async_calls.append(kwargs["messages"])
1519-
mock_resp = MagicMock()
1520-
mock_resp.choices = [MagicMock()]
1521-
mock_resp.choices[0].message.content = concept_response
1522-
mock_resp.usage = MagicMock(prompt_tokens=1, completion_tokens=1)
1523-
mock_resp.usage.prompt_tokens_details = None
1524-
return mock_resp
1515+
return [_mock_response(concept_response)]
15251516

15261517
with patch("openkb.agent.compiler.litellm") as mock_litellm:
15271518
mock_litellm.completion = MagicMock(side_effect=sync_side_effect)
@@ -1586,15 +1577,9 @@ async def test_long_doc_marks_doc_message(self, tmp_path):
15861577

15871578
def sync_side_effect(*args, **kwargs):
15881579
captured.append(kwargs["messages"])
1589-
mock_resp = MagicMock()
1590-
mock_resp.choices = [MagicMock()]
15911580
# First call: overview (plain text); second: plan (JSON).
1592-
mock_resp.choices[0].message.content = (
1593-
"Overview text" if len(captured) == 1 else plan_response
1594-
)
1595-
mock_resp.usage = MagicMock(prompt_tokens=1, completion_tokens=1)
1596-
mock_resp.usage.prompt_tokens_details = None
1597-
return mock_resp
1581+
content = "Overview text" if len(captured) == 1 else plan_response
1582+
return [_mock_response(content)]
15981583

15991584
with patch("openkb.agent.compiler.litellm") as mock_litellm:
16001585
mock_litellm.completion = MagicMock(side_effect=sync_side_effect)
@@ -1726,16 +1711,9 @@ async def test_create_and_update_flow(self, tmp_path):
17261711
async def ordered_acompletion(*args, **kwargs):
17271712
idx = call_order["n"]
17281713
call_order["n"] += 1
1729-
mock_resp = MagicMock()
1730-
mock_resp.choices = [MagicMock()]
17311714
# create tasks come first, then update tasks
1732-
if idx == 0:
1733-
mock_resp.choices[0].message.content = create_page_response
1734-
else:
1735-
mock_resp.choices[0].message.content = update_page_response
1736-
mock_resp.usage = MagicMock(prompt_tokens=100, completion_tokens=50)
1737-
mock_resp.usage.prompt_tokens_details = None
1738-
return mock_resp
1715+
content = create_page_response if idx == 0 else update_page_response
1716+
return [_mock_response(content)]
17391717

17401718
with patch("openkb.agent.compiler.litellm") as mock_litellm:
17411719
mock_litellm.completion = MagicMock(side_effect=_mock_completion([plan_response]))
@@ -1823,13 +1801,7 @@ async def test_truncated_update_preserves_existing_page(self, tmp_path):
18231801
)
18241802

18251803
async def truncated_acompletion(*args, **kwargs):
1826-
mock_resp = MagicMock()
1827-
mock_resp.choices = [MagicMock()]
1828-
mock_resp.choices[0].message.content = truncated_page
1829-
mock_resp.choices[0].finish_reason = "length"
1830-
mock_resp.usage = MagicMock(prompt_tokens=100, completion_tokens=50)
1831-
mock_resp.usage.prompt_tokens_details = None
1832-
return mock_resp
1804+
return [_mock_response(truncated_page, finish_reason="length")]
18331805

18341806
with patch("openkb.agent.compiler.litellm") as mock_litellm:
18351807
mock_litellm.completion = MagicMock(side_effect=_mock_completion([plan_response]))
@@ -1859,13 +1831,7 @@ async def test_truncated_create_skips_partial_page(self, tmp_path):
18591831
truncated_page = json.dumps({"brief": "x", "content": "# Ghost\n\nPartial"})
18601832

18611833
async def truncated_acompletion(*args, **kwargs):
1862-
mock_resp = MagicMock()
1863-
mock_resp.choices = [MagicMock()]
1864-
mock_resp.choices[0].message.content = truncated_page
1865-
mock_resp.choices[0].finish_reason = "length"
1866-
mock_resp.usage = MagicMock(prompt_tokens=100, completion_tokens=50)
1867-
mock_resp.usage.prompt_tokens_details = None
1868-
return mock_resp
1834+
return [_mock_response(truncated_page, finish_reason="length")]
18691835

18701836
with patch("openkb.agent.compiler.litellm") as mock_litellm:
18711837
mock_litellm.completion = MagicMock(side_effect=_mock_completion([plan_response]))
@@ -1928,13 +1894,7 @@ async def test_truncated_entity_update_preserves_existing_page(self, tmp_path):
19281894
)
19291895

19301896
async def truncated_acompletion(*args, **kwargs):
1931-
mock_resp = MagicMock()
1932-
mock_resp.choices = [MagicMock()]
1933-
mock_resp.choices[0].message.content = truncated_page
1934-
mock_resp.choices[0].finish_reason = "length"
1935-
mock_resp.usage = MagicMock(prompt_tokens=100, completion_tokens=50)
1936-
mock_resp.usage.prompt_tokens_details = None
1937-
return mock_resp
1897+
return [_mock_response(truncated_page, finish_reason="length")]
19381898

19391899
with patch("openkb.agent.compiler.litellm") as mock_litellm:
19401900
mock_litellm.completion = MagicMock(side_effect=_mock_completion([plan_response]))

‎tests/test_llm_timeout.py‎

Lines changed: 13 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,13 @@
1717

1818

1919
def _fake_response():
20+
"""A fake, already-complete LLM response (single-chunk stream).
21+
22+
See ``openkb.agent.compiler._merge_stream_chunks``: a chunk exposing
23+
``.message`` (as this one does) is treated as already-complete and used
24+
as-is, so callers of ``litellm.completion``/``acompletion`` with
25+
``stream=True`` can be mocked to just return a one-item list.
26+
"""
2027
choice = MagicMock()
2128
choice.message.content = "ok"
2229
choice.finish_reason = "stop"
@@ -28,7 +35,7 @@ def _fake_response():
2835
def test_llm_call_forwards_configured_timeout():
2936
set_timeout(1200.0)
3037
with patch(
31-
"openkb.agent.compiler.litellm.completion", return_value=_fake_response()
38+
"openkb.agent.compiler.litellm.completion", return_value=[_fake_response()]
3239
) as completion:
3340
_llm_call("gpt-4o", [{"role": "user", "content": "hi"}], "step")
3441
assert completion.call_args.kwargs["timeout"] == 1200.0
@@ -37,7 +44,7 @@ def test_llm_call_forwards_configured_timeout():
3744
def test_llm_call_omits_timeout_when_unset():
3845
set_timeout(None)
3946
with patch(
40-
"openkb.agent.compiler.litellm.completion", return_value=_fake_response()
47+
"openkb.agent.compiler.litellm.completion", return_value=[_fake_response()]
4148
) as completion:
4249
_llm_call("gpt-4o", [{"role": "user", "content": "hi"}], "step")
4350
assert "timeout" not in completion.call_args.kwargs
@@ -47,7 +54,7 @@ def test_llm_call_does_not_override_explicit_timeout():
4754
# An explicit per-call timeout kwarg wins over the configured default.
4855
set_timeout(1200.0)
4956
with patch(
50-
"openkb.agent.compiler.litellm.completion", return_value=_fake_response()
57+
"openkb.agent.compiler.litellm.completion", return_value=[_fake_response()]
5158
) as completion:
5259
_llm_call("gpt-4o", [{"role": "user", "content": "hi"}], "step", timeout=30)
5360
assert completion.call_args.kwargs["timeout"] == 30
@@ -58,7 +65,7 @@ def test_llm_call_async_forwards_configured_timeout():
5865
with patch(
5966
"openkb.agent.compiler.litellm.acompletion",
6067
new_callable=AsyncMock,
61-
return_value=_fake_response(),
68+
return_value=[_fake_response()],
6269
) as acompletion:
6370
asyncio.run(_llm_call_async("gpt-4o", [{"role": "user", "content": "hi"}], "step"))
6471
assert acompletion.call_args.kwargs["timeout"] == 900.0
@@ -69,7 +76,7 @@ def test_llm_call_async_omits_timeout_when_unset():
6976
with patch(
7077
"openkb.agent.compiler.litellm.acompletion",
7178
new_callable=AsyncMock,
72-
return_value=_fake_response(),
79+
return_value=[_fake_response()],
7380
) as acompletion:
7481
asyncio.run(_llm_call_async("gpt-4o", [{"role": "user", "content": "hi"}], "step"))
7582
assert "timeout" not in acompletion.call_args.kwargs
@@ -80,7 +87,7 @@ def test_llm_call_async_does_not_override_explicit_timeout():
8087
with patch(
8188
"openkb.agent.compiler.litellm.acompletion",
8289
new_callable=AsyncMock,
83-
return_value=_fake_response(),
90+
return_value=[_fake_response()],
8491
) as acompletion:
8592
asyncio.run(
8693
_llm_call_async("gpt-4o", [{"role": "user", "content": "hi"}], "step", timeout=30)

0 commit comments

Comments
 (0)