🐛(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:
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
"""
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user