🐛(back) stream tool responses to prevent too call timeouts

Implement sync/sync utilities that inject
keepalive messages at regular intervals during stream pauses,
preventing proxy timeouts on long-running operations like
document(s) summarization.

Keepalive messages maintain active connections while tools execute,
eliminating forced conversation restarts.

Signed-off-by: Laurent Paoletti <[email protected]>
This commit is contained in:
Laurent Paoletti
2026-01-17 13:50:35 +01:00
parent e7d76e4477
commit ddfc86a88f
9 changed files with 516 additions and 37 deletions
+1
View File
@@ -19,6 +19,7 @@ and this project adheres to
- 🐛(e2e) fix test-e2e-chromium
- 🐛(back) fix system prompt compatibility with self-hosted models #200
- ⚰️(back) remove dead code and unused files
- 🐛(back) prevent tool call timeouts
### Removed
+2 -2
View File
@@ -76,7 +76,7 @@ from chat.tools.document_generic_search_rag import add_document_rag_search_tool_
from chat.tools.document_search_rag import add_document_rag_search_tool
from chat.tools.document_summarize import document_summarize
from chat.vercel_ai_sdk.core import events_v4, events_v5
from chat.vercel_ai_sdk.encoder import EventEncoder
from chat.vercel_ai_sdk.encoder import CURRENT_EVENT_ENCODER_VERSION, EventEncoder
# Keep at the top of the file to avoid mocking issues
document_store_backend = import_string(settings.RAG_DOCUMENT_SEARCH_BACKEND)
@@ -122,7 +122,7 @@ class AIAgentService: # pylint: disable=too-many-instance-attributes
self._langfuse_available = settings.LANGFUSE_ENABLED
self._store_analytics = self._langfuse_available and user.allow_conversation_analytics
self.event_encoder = EventEncoder("v4") # Always use v4 for now
self.event_encoder = EventEncoder(CURRENT_EVENT_ENCODER_VERSION) # We use v4 for now
self._support_streaming = True
if (streaming := get_model_configuration(self.model_hrid).supports_streaming) is not None:
+171
View File
@@ -0,0 +1,171 @@
"""Helpers to prevent proxy timeouts during long-running stream operations.
This module provides utilities to wrap synchronous and asynchronous iterators
with keepalive messages. When a stream pauses for longer than the specified
interval, keepalive messages are injected to prevent proxy/gateway
timeouts while waiting for the stream data.
"""
import asyncio
import logging
import queue
import threading
import time
from typing import AsyncIterator, Iterator
from django.conf import settings
from .vercel_ai_sdk.core.events_v4 import DataPart as V4DataPart
from .vercel_ai_sdk.core.events_v5 import DataPart as V5DataPart
from .vercel_ai_sdk.encoder import (
CURRENT_EVENT_ENCODER_VERSION,
EventEncoder,
EventEncoderVersion,
)
logger = logging.getLogger(__name__)
def get_keepalive_message() -> str:
"""Generate a keepalive message based on encoder/SDK version."""
if CURRENT_EVENT_ENCODER_VERSION == EventEncoderVersion.V4:
event = V4DataPart(data=[{"status": "WAITING"}])
else:
event = V5DataPart(data={"status": "WAITING"})
encoder = EventEncoder(CURRENT_EVENT_ENCODER_VERSION)
return encoder.encode(event)
async def stream_with_keepalive_async(
stream: AsyncIterator[str],
) -> AsyncIterator[str]:
"""Wrap an async iterator to emit keepalive during long pauses.
Args:
stream: The async iterator to wrap
Yields:
Items from the original stream, plus keepalive messages during pauses
Raises:
Any exception raised by the original stream
"""
q: asyncio.Queue = asyncio.Queue()
finished = asyncio.Event()
keepalive_message = get_keepalive_message()
async def producer():
"""Background task that consumes the original stream into a queue."""
try:
async for stream_item in stream:
await q.put(stream_item)
except Exception as exc: # pylint: disable=broad-except #noqa: BLE001
# Pass exceptions through the queue so the consumer can re-raise them.
# This ensures errors aren't silently swallowed.
await q.put(exc)
finally:
finished.set()
await q.put(None) # Sentinel to signal completion
producer_task = asyncio.create_task(producer())
try:
while True:
try:
item = await asyncio.wait_for(q.get(), timeout=settings.KEEPALIVE_INTERVAL)
if item is None:
break
if isinstance(item, Exception):
raise item
yield item
except asyncio.TimeoutError:
# No data received within interval
if finished.is_set():
# Producer is done, queue is empty (else we would not have timed out)
break
logger.debug("Send keepalive")
yield keepalive_message
finally:
# Cleanup
producer_task.cancel()
try:
await producer_task
except asyncio.CancelledError:
pass
def get_current_time() -> float:
"""Get current monotonic time, avoiding freezegun interferences.
Returns time.monotonic() which:
- Is NOT affected by freezegun's @freeze_time decorator (unlike time.time())
- Prevents issues where frozen time in main thread differs from real time in
spawned threads, causing incorrect keepalive interval computation
- Is the best clock for measuring time intervals
Wrapped in a function to ease mocking in tests.
Returns:
float: Monotonic time in seconds since an arbitrary reference point
"""
return time.monotonic()
def stream_with_keepalive_sync(stream: Iterator[str]) -> Iterator[str]:
"""Wraps a synchronous stream with keepalive messages."""
q: queue.Queue = queue.Queue()
stream_done = threading.Event()
keepalive_message = get_keepalive_message()
# Mutable container so threads can read/write shared timestamp
last_yield_time = [get_current_time()]
def consume_stream():
"""Read from source stream and forward chunks to queue."""
try:
for chunk in stream:
if stream_done.is_set():
return # early exit
q.put(chunk, timeout=1) # Arbitrary timeout prevents blocking forever
# pylint: disable=broad-exception-caught
except Exception as e:
logger.exception("Error in stream consumption")
q.put(e)
finally:
stream_done.set()
def send_keepalives():
"""Inject keepalive messages when idle too long.
Uses get_current_time() (time.monotonic) instead of time.time()
to avoid issues with freezegun in tests.
"""
while not stream_done.is_set():
# Sleep before checking to give main loop time to process and update timestamp
time.sleep(0.5) # let main loop process first, empiric value
if get_current_time() - last_yield_time[0] >= settings.KEEPALIVE_INTERVAL:
try:
q.put(keepalive_message, timeout=0.1)
except queue.Full:
pass
for target in (consume_stream, send_keepalives):
threading.Thread(target=target, daemon=True).start()
try:
# Continue while stream is active or queue has still items
while not stream_done.is_set() or not q.empty():
try:
item = q.get(timeout=1) # short timeout, avoid blocking and stay responsive
except queue.Empty:
continue
# Re-raise from consume_stream
if isinstance(item, Exception):
raise item
yield item
last_yield_time[0] = get_current_time()
finally:
# Signal threads to stop
stream_done.set()
@@ -1,5 +1,6 @@
"""Common test fixtures for chat conversation endpoint tests."""
import asyncio
import json
from django.utils import timezone
@@ -10,15 +11,9 @@ import respx
from freezegun import freeze_time
@pytest.fixture(name="mock_openai_stream")
@freeze_time("2025-07-25T10:36:35.297675Z")
def fixture_mock_openai_stream():
"""
Fixture to mock the OpenAI stream response.
See https://platform.openai.com/docs/api-reference/chat-streaming/streaming
"""
openai_stream = (
def _create_openai_stream_data():
"""Helper to create OpenAI stream data."""
return (
"data: "
+ json.dumps(
{
@@ -59,15 +54,43 @@ def fixture_mock_openai_stream():
"data: [DONE]\n\n"
)
async def mock_stream():
for line in openai_stream.splitlines(keepends=True):
yield line.encode()
route = respx.post("https://www.external-ai-service.com/chat/completions").mock(
def _create_mock_openai_route(with_delays: bool = False, delay_seconds: float = 1.0):
"""Create a mock OpenAI stream route with optional delays."""
openai_stream = _create_openai_stream_data()
async def mock_stream():
lines = openai_stream.splitlines(keepends=True)
for i, line in enumerate(lines):
yield line.encode()
if with_delays and i == 1:
# Delay after second line to trigger keepalive during streaming
await asyncio.sleep(delay_seconds)
return respx.post("https://www.external-ai-service.com/chat/completions").mock(
return_value=httpx.Response(200, stream=mock_stream())
)
return route
@pytest.fixture(name="mock_openai_stream")
@freeze_time("2025-07-25T10:36:35.297675Z")
def fixture_mock_openai_stream():
"""
Fixture to mock the OpenAI stream response (no delays).
See https://platform.openai.com/docs/api-reference/chat-streaming/streaming
"""
return _create_mock_openai_route(with_delays=False)
@pytest.fixture(name="mock_openai_stream_slow")
def fixture_mock_openai_stream_slow():
"""
Fixture to mock the OpenAI stream response with delays to trigger keepalives.
No @freeze_time decorator because asyncio.sleep() needs real time to work properly.
"""
return _create_mock_openai_route(with_delays=True, delay_seconds=1.0)
@pytest.fixture(name="mock_openai_no_stream")
@@ -3,6 +3,7 @@
import json
import logging
from unittest.mock import ANY, patch
from django.utils import timezone
@@ -221,6 +222,133 @@ def test_post_conversation_data_protocol(api_client, mock_openai_stream):
]
@freeze_time("2025-07-25T10:36:35.297675Z")
@respx.mock
@patch("chat.keepalive.get_current_time")
def test_post_conversation_data_protocol_triggers_keepalives(
mock_time, api_client, mock_openai_stream
):
"""Test streaming response contains keepalive messages"""
chat_conversation = ChatConversationFactory(owner__language="en-us")
mock_time.side_effect = [float(i * 60) for i in range(10)]
url = f"/api/v1.0/chats/{chat_conversation.pk}/conversation/?protocol=data"
data = {
"messages": [
{
"id": "yuPoOuBkKA4FnKvk",
"role": "user",
"parts": [{"text": "Hello", "type": "text"}],
"content": "Hello",
"createdAt": "2025-07-03T15:22:17.105Z",
}
]
}
api_client.force_login(chat_conversation.owner)
response = api_client.post(url, data, format="json")
assert response.status_code == status.HTTP_200_OK
assert response.get("Content-Type") == "text/event-stream"
assert response.get("x-vercel-ai-data-stream") == "v1"
assert response.streaming
# Wait for the streaming content to be fully received
response_content = b"".join(response.streaming_content).decode("utf-8")
# Replace UUIDs with placeholders for assertion
response_content = replace_uuids_with_placeholder(response_content)
assert response_content == (
'0:"Hello"\n'
'0:" there"\n'
'f:{"messageId":"<mocked_uuid>"}\n'
'd:{"finishReason":"stop","usage":{"promptTokens":0,"completionTokens":0}}\n'
'2:[{"status": "WAITING"}]\n'
)
assert mock_openai_stream.called
chat_conversation.refresh_from_db()
assert chat_conversation.ui_messages == [
{
"content": "Hello",
"createdAt": "2025-07-03T15:22:17.105Z",
"id": "yuPoOuBkKA4FnKvk",
"parts": [{"text": "Hello", "type": "text"}],
"role": "user",
}
]
assert len(chat_conversation.messages) == 2
assert chat_conversation.messages[0].id == IsUUID(4)
assert chat_conversation.messages[0] == UIMessage(
id=chat_conversation.messages[0].id, # don't test the message ID here
createdAt=timezone.now(), # Mocked timestamp
content="Hello",
reasoning=None,
experimental_attachments=None,
role="user",
annotations=None,
toolInvocations=None,
parts=[TextUIPart(type="text", text="Hello")],
)
assert chat_conversation.messages[1].id == IsUUID(4)
assert chat_conversation.messages[1] == UIMessage(
id=chat_conversation.messages[1].id, # don't test the message ID here
createdAt=timezone.now(), # Mocked timestamp
content="Hello there",
reasoning=None,
experimental_attachments=None,
role="assistant",
annotations=None,
toolInvocations=None,
parts=[TextUIPart(type="text", text="Hello there")],
)
_run_id = chat_conversation.pydantic_messages[0]["run_id"]
assert chat_conversation.pydantic_messages == [
{
"instructions": (
"You are a helpful test assistant :)\n\n"
"Today is Friday 25/07/2025.\n\n"
"Answer in english."
),
"kind": "request",
"parts": [
{
"content": ["Hello"],
"part_kind": "user-prompt",
"timestamp": "2025-07-25T10:36:35.297675Z",
},
],
"run_id": _run_id,
},
{
"finish_reason": "stop",
"kind": "response",
"model_name": "test-model",
"parts": [{"content": "Hello there", "id": None, "part_kind": "text"}],
"provider_details": {"finish_reason": "stop"},
"provider_name": "openai",
"provider_response_id": "chatcmpl-1234567890",
"timestamp": "2025-07-25T10:36:35.297675Z",
"usage": {
"cache_audio_read_tokens": 0,
"cache_read_tokens": 0,
"cache_write_tokens": 0,
"details": {},
"input_audio_tokens": 0,
"input_tokens": 0,
"output_audio_tokens": 0,
"output_tokens": 0,
},
"run_id": _run_id,
},
]
@freeze_time("2025-07-25T10:36:35.297675Z")
@respx.mock
def test_post_conversation_text_protocol(api_client, mock_openai_stream):
@@ -1344,3 +1472,143 @@ async def test_post_conversation_async(api_client, mock_openai_stream, monkeypat
"run_id": _run_id,
},
]
@freeze_time("2025-07-25T10:36:35.297675Z", tick=True)
@respx.mock
@pytest.mark.asyncio
async def test_post_conversation_async_triggers_keepalive(
api_client, mock_openai_stream_slow, monkeypatch, caplog, settings
):
"""Test posting messages to a conversation using the 'data' protocol."""
monkeypatch.setenv("PYTHON_SERVER_MODE", "async")
settings.KEEPALIVE_INTERVAL = 1 # s
chat_conversation = await sync_to_async(ChatConversationFactory)(owner__language="en-us")
url = f"/api/v1.0/chats/{chat_conversation.pk}/conversation/?protocol=data"
data = {
"messages": [
{
"id": "yuPoOuBkKA4FnKvk",
"role": "user",
"parts": [{"text": "Hello", "type": "text"}],
"content": "Hello",
"createdAt": "2025-07-03T15:22:17.105Z",
}
]
}
await api_client.aforce_login(chat_conversation.owner)
caplog.clear()
caplog.set_level(level=logging.DEBUG, logger="chat.views")
response = await sync_to_async(api_client.post)(url, data, format="json") # client is sync
assert response.status_code == status.HTTP_200_OK
assert response.get("Content-Type") == "text/event-stream"
assert response.get("x-vercel-ai-data-stream") == "v1"
assert response.streaming
assert "Using ASYNC streaming for chat conversation" in caplog.text
# Wait for the streaming content to be fully received => async iterator -> list
# This fails it the streaming is not an async generator
response_content = b"".join([content async for content in response.streaming_content]).decode(
"utf-8"
)
# Replace UUIDs with placeholders for assertion
response_content = replace_uuids_with_placeholder(response_content)
assert response_content == (
'0:"Hello"\n'
'2:[{"status": "WAITING"}]\n'
'0:" there"\n'
'f:{"messageId":"<mocked_uuid>"}\n'
'd:{"finishReason":"stop","usage":{"promptTokens":0,"completionTokens":0}}\n'
)
assert mock_openai_stream_slow.called
await chat_conversation.arefresh_from_db()
assert chat_conversation.ui_messages == [
{
"content": "Hello",
"createdAt": "2025-07-03T15:22:17.105Z",
"id": "yuPoOuBkKA4FnKvk",
"parts": [{"text": "Hello", "type": "text"}],
"role": "user",
}
]
assert len(chat_conversation.messages) == 2
assert chat_conversation.messages[0].id == IsUUID(4)
assert chat_conversation.messages[0] == UIMessage(
id=chat_conversation.messages[0].id, # don't test the message ID here
createdAt=chat_conversation.messages[0].createdAt, # Mocked timestamp
content="Hello",
reasoning=None,
experimental_attachments=None,
role="user",
annotations=None,
toolInvocations=None,
parts=[TextUIPart(type="text", text="Hello")],
)
assert chat_conversation.messages[1].id == IsUUID(4)
assert chat_conversation.messages[1] == UIMessage(
id=chat_conversation.messages[1].id, # don't test the message ID here
createdAt=chat_conversation.messages[1].createdAt, # Mocked timestamp
content="Hello there",
reasoning=None,
experimental_attachments=None,
role="assistant",
annotations=None,
toolInvocations=None,
parts=[TextUIPart(type="text", text="Hello there")],
)
_run_id = chat_conversation.pydantic_messages[0]["run_id"]
# using ANY because time is not frozen in this api mock
assert chat_conversation.pydantic_messages == [
{
"instructions": (
"You are a helpful test assistant :)\n\n"
"Today is Friday 25/07/2025.\n\nAnswer in english."
),
"kind": "request",
"parts": [
{
"content": ["Hello"],
"part_kind": "user-prompt",
"timestamp": ANY,
},
],
"run_id": _run_id,
},
{
"finish_reason": "stop",
"kind": "response",
"model_name": "test-model",
"parts": [{"content": "Hello there", "id": None, "part_kind": "text"}],
"provider_details": {"finish_reason": "stop"},
"provider_name": "openai",
"provider_response_id": "chatcmpl-1234567890",
"timestamp": ANY,
"usage": {
"cache_audio_read_tokens": 0,
"cache_read_tokens": 0,
"cache_write_tokens": 0,
"details": {},
"input_audio_tokens": 0,
"input_tokens": 0,
"output_audio_tokens": 0,
"output_tokens": 0,
},
"run_id": _run_id,
},
]
@@ -2,6 +2,6 @@
This module contains the EventEncoder class.
"""
from .encoder import EventEncoder
from .encoder import CURRENT_EVENT_ENCODER_VERSION, EventEncoder, EventEncoderVersion
__all__ = ["EventEncoder"]
__all__ = ["EventEncoder", "CURRENT_EVENT_ENCODER_VERSION", "EventEncoderVersion"]
@@ -1,6 +1,7 @@
"""Event Encoder for Vercel AI SDK"""
from typing import Literal, Union
from enum import Enum
from typing import Union
from ..core.events_v4 import BaseEvent as V4BaseEvent
from ..core.events_v4 import TextPart
@@ -8,16 +9,26 @@ from ..core.events_v5 import BaseEvent as V5BaseEvent
from ..core.events_v5 import TextDeltaEvent
class EventEncoderVersion(str, Enum):
"""Enumeration of supported event encoder versions."""
V4 = "v4"
V5 = "v5"
CURRENT_EVENT_ENCODER_VERSION = EventEncoderVersion.V4 # used encoder version
class EventEncoder:
"""
Encodes events for the Vercel AI SDK based on the specified version.
"""
def __init__(self, version: Literal["v4", "v5"] = None):
def __init__(self, version: EventEncoderVersion):
"""
Initializes the EventEncoder with the specified version.
"""
if version not in ["v4", "v5"]:
if version not in [EventEncoderVersion.V4, EventEncoderVersion.V5]:
raise ValueError("Unsupported version. Supported versions are 'v4' and 'v5'.")
self.version = version
@@ -28,7 +39,7 @@ class EventEncoder:
"""
return "text/event-stream"
def encode(self, event: Union[V5BaseEvent, V5BaseEvent]) -> str | None:
def encode(self, event: Union[V4BaseEvent, V5BaseEvent]) -> str | None:
"""
Encodes an event based on the version.
@@ -38,15 +49,15 @@ class EventEncoder:
str | None: The encoded event as a string,
or None if the event type is not adapted to the SDK version.
"""
if self.version == "v4" and isinstance(event, V4BaseEvent):
if self.version == EventEncoderVersion.V4 and isinstance(event, V4BaseEvent):
return self._encode_v4_streaming(event)
if self.version == "v5" and isinstance(event, V5BaseEvent):
if self.version == EventEncoderVersion.V5 and isinstance(event, V5BaseEvent):
return self._encode_sse(event)
return None
def encode_text(self, event: Union[V5BaseEvent, V5BaseEvent]) -> str | None:
def encode_text(self, event: Union[V4BaseEvent, V5BaseEvent]) -> str | None:
"""
Encodes an event based on the version.
@@ -56,10 +67,10 @@ class EventEncoder:
str | None: The encoded event as a string,
or None if the event type is not adapted to the SDK version.
"""
if self.version == "v4" and isinstance(event, TextPart):
if self.version == EventEncoderVersion.V4 and isinstance(event, TextPart):
return event.text
if self.version == "v5" and isinstance(event, TextDeltaEvent):
if self.version == EventEncoderVersion.V5 and isinstance(event, TextDeltaEvent):
return event.delta
return None
@@ -70,7 +81,7 @@ class EventEncoder:
"""
return f"{event.type}:{event.model_dump_json(by_alias=True, exclude={'type'})}\n"
def _encode_sse(self, event: Union[V5BaseEvent, V5BaseEvent]) -> str:
def _encode_sse(self, event: Union[V4BaseEvent, V5BaseEvent]) -> str:
"""
Encodes an event into an SSE string.
"""
+8 -9
View File
@@ -26,6 +26,7 @@ from core.filters import remove_accents
from activation_codes.permissions import IsActivatedUser
from chat import models, serializers
from chat.clients.pydantic_ai import AIAgentService
from chat.keepalive import stream_with_keepalive_async, stream_with_keepalive_sync
from chat.serializers import ChatConversationRequestSerializer
logger = logging.getLogger(__name__)
@@ -188,29 +189,28 @@ class ChatViewSet( # pylint: disable=too-many-ancestors, abstract-method
if is_async_mode:
logger.debug("Using ASYNC streaming for chat conversation.")
if protocol == "data":
streaming_content = ai_service.stream_data_async(
base_stream = ai_service.stream_data_async(
messages, force_web_search=force_web_search
)
else: # Default to 'text' protocol
streaming_content = ai_service.stream_text_async(
base_stream = ai_service.stream_text_async(
messages, force_web_search=force_web_search
)
streaming_content = stream_with_keepalive_async(base_stream)
else:
logger.debug("Using SYNC streaming for chat conversation.")
if protocol == "data":
streaming_content = ai_service.stream_data(
messages, force_web_search=force_web_search
)
base_stream = ai_service.stream_data(messages, force_web_search=force_web_search)
else: # Default to 'text' protocol
streaming_content = ai_service.stream_text(
messages, force_web_search=force_web_search
)
base_stream = ai_service.stream_text(messages, force_web_search=force_web_search)
streaming_content = stream_with_keepalive_sync(base_stream)
response = StreamingHttpResponse(
streaming_content,
content_type="text/event-stream",
headers={
"x-vercel-ai-data-stream": "v1", # This header is used for Vercel AI streaming,
"X-Accel-Buffering": "no", # Prevent nginx buffering
},
)
return response
@@ -371,7 +371,6 @@ class ChatConversationAttachmentViewSet(
owner=self.request.user,
).exists():
raise Http404
file_name = serializer.validated_data["file_name"]
extension = file_name.rpartition(".")[-1] if "." in file_name else None
+6
View File
@@ -919,6 +919,12 @@ USER QUESTION:
environ_prefix=None,
)
# Default keepalive interval: 55s (safely below typical 60s proxy timeouts)
# Prevents connection drops during long stream pauses while providing 5s safety margin.
KEEPALIVE_INTERVAL = values.PositiveIntegerValue(
default=55, environ_name="KEEPALIVE_INTERVAL", environ_prefix=None
)
# pylint: disable=invalid-name
@property
def ENVIRONMENT(self):