Skip to content

Commit 77bcbe7

Browse files
feat(workflows): raise StreamDisconnectedError on SSE error frames
Adds an AfterSuccess hook that converts a workflow SSE `event: error` frame into a raised StreamDisconnectedError (reason + error), so consumers use try/except around stream iteration instead of inspecting each event.
1 parent 7e7c967 commit 77bcbe7

5 files changed

Lines changed: 374 additions & 0 deletions

File tree

‎src/mistralai/client/_hooks/registration.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
from .tracing import TracingHook
55
from .types import Hooks
66
from .workflow_encoding_hook import WorkflowEncodingHook
7+
from mistralai.extra.workflows.stream_error_hook import WorkflowStreamErrorHook
78

89
# This file is only ever generated once on the first generation and then is free to be modified.
910
# Any hooks you wish to add should be registered in the init_hooks function. Feel free to define them
@@ -26,3 +27,4 @@ def init_hooks(hooks: Hooks):
2627
hooks.register_after_error_hook(tracing_hook)
2728
hooks.register_before_request_hook(workflow_encoding_hook)
2829
hooks.register_after_success_hook(workflow_encoding_hook)
30+
hooks.register_after_success_hook(WorkflowStreamErrorHook())
Lines changed: 154 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,154 @@
1+
import httpx
2+
import pytest
3+
from httpx._types import AsyncByteStream, SyncByteStream
4+
5+
from mistralai.client import Mistral
6+
from mistralai.client._hooks.types import AfterSuccessContext, HookContext
7+
from mistralai.extra.workflows.errors import StreamDisconnectedError
8+
from mistralai.extra.workflows.stream_error_hook import WorkflowStreamErrorHook
9+
10+
STREAM_OPERATION_ID = "get_stream_events_v1_workflows_events_stream_get"
11+
NON_STREAM_OPERATION_ID = "chat_completion_v1_chat_completions_post"
12+
13+
GOOD_FRAME = b'event: workflow.event\ndata: {"attributes": {}}\n\n'
14+
ERROR_FRAME = b'event: error\ndata: {"error": "boom", "reason": "read_error"}\n\n'
15+
16+
17+
class _SyncSource(SyncByteStream):
18+
def __init__(self, chunks: list[bytes]) -> None:
19+
self._chunks = chunks
20+
21+
def __iter__(self):
22+
yield from self._chunks
23+
24+
def close(self) -> None:
25+
pass
26+
27+
28+
class _AsyncSource(AsyncByteStream):
29+
def __init__(self, chunks: list[bytes]) -> None:
30+
self._iter = iter(chunks)
31+
32+
def __aiter__(self) -> "_AsyncSource":
33+
return self
34+
35+
async def __anext__(self) -> bytes:
36+
try:
37+
return next(self._iter)
38+
except StopIteration:
39+
raise StopAsyncIteration
40+
41+
async def aclose(self) -> None:
42+
pass
43+
44+
45+
def _hook_ctx(operation_id: str) -> AfterSuccessContext:
46+
client = Mistral(api_key="test-key")
47+
return AfterSuccessContext(
48+
HookContext(
49+
config=client.sdk_configuration,
50+
base_url="https://api.example.com",
51+
operation_id=operation_id,
52+
oauth2_scopes=[],
53+
security_source=None,
54+
)
55+
)
56+
57+
58+
def _sse_response(source, *, content_type: str = "text/event-stream") -> httpx.Response:
59+
return httpx.Response(
60+
status_code=200,
61+
headers={"content-type": content_type},
62+
stream=source,
63+
request=httpx.Request(
64+
"GET", "https://api.example.com/v1/workflows/events/stream"
65+
),
66+
)
67+
68+
69+
def test_hook_raises_stream_disconnected_error_and_passes_prior_events_through():
70+
response = _sse_response(_SyncSource([GOOD_FRAME, ERROR_FRAME]))
71+
result = WorkflowStreamErrorHook().after_success(
72+
_hook_ctx(STREAM_OPERATION_ID), response
73+
)
74+
assert isinstance(result, httpx.Response)
75+
76+
collected: list[bytes] = []
77+
with pytest.raises(StreamDisconnectedError) as exc_info:
78+
for chunk in result.iter_bytes():
79+
collected.append(chunk)
80+
81+
assert exc_info.value.reason == "read_error"
82+
assert exc_info.value.error == "boom"
83+
assert b"workflow.event" in b"".join(collected)
84+
85+
86+
@pytest.mark.asyncio
87+
async def test_hook_raises_stream_disconnected_error_on_async_stream():
88+
response = _sse_response(_AsyncSource([GOOD_FRAME, ERROR_FRAME]))
89+
result = WorkflowStreamErrorHook().after_success(
90+
_hook_ctx(STREAM_OPERATION_ID), response
91+
)
92+
assert isinstance(result, httpx.Response)
93+
94+
collected: list[bytes] = []
95+
with pytest.raises(StreamDisconnectedError) as exc_info:
96+
async for chunk in result.aiter_bytes():
97+
collected.append(chunk)
98+
99+
assert exc_info.value.reason == "read_error"
100+
assert exc_info.value.error == "boom"
101+
assert b"workflow.event" in b"".join(collected)
102+
103+
104+
def test_hook_detects_error_frame_split_across_chunks():
105+
chunks = [
106+
b"event: er",
107+
b'ror\ndata: {"error": "x", "reason": "internal_error"}\n\n',
108+
]
109+
response = _sse_response(_SyncSource(chunks))
110+
result = WorkflowStreamErrorHook().after_success(
111+
_hook_ctx(STREAM_OPERATION_ID), response
112+
)
113+
114+
with pytest.raises(StreamDisconnectedError) as exc_info:
115+
list(result.iter_bytes())
116+
117+
assert exc_info.value.reason == "internal_error"
118+
assert exc_info.value.error == "x"
119+
120+
121+
def test_hook_defaults_reason_when_missing_or_invalid():
122+
frame = b'event: error\ndata: {"error": "no reason given"}\n\n'
123+
response = _sse_response(_SyncSource([frame]))
124+
result = WorkflowStreamErrorHook().after_success(
125+
_hook_ctx(STREAM_OPERATION_ID), response
126+
)
127+
128+
with pytest.raises(StreamDisconnectedError) as exc_info:
129+
list(result.iter_bytes())
130+
131+
assert exc_info.value.reason == "stream_error"
132+
assert exc_info.value.error == "no reason given"
133+
134+
135+
def test_hook_passes_normal_stream_through_without_raising():
136+
response = _sse_response(_SyncSource([GOOD_FRAME, GOOD_FRAME]))
137+
result = WorkflowStreamErrorHook().after_success(
138+
_hook_ctx(STREAM_OPERATION_ID), response
139+
)
140+
141+
body = b"".join(result.iter_bytes())
142+
assert body.count(b"workflow.event") == 2
143+
144+
145+
def test_hook_ignores_non_stream_operations():
146+
source = _SyncSource([ERROR_FRAME])
147+
response = _sse_response(source)
148+
result = WorkflowStreamErrorHook().after_success(
149+
_hook_ctx(NON_STREAM_OPERATION_ID), response
150+
)
151+
152+
# Response is returned untouched: same object, original stream not wrapped.
153+
assert result is response
154+
assert response.stream is source

‎src/mistralai/extra/workflows/__init__.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,10 @@
2222
configure_workflow_encoding,
2323
generate_two_part_id,
2424
)
25+
from .errors import (
26+
StreamDisconnectReason,
27+
StreamDisconnectedError,
28+
)
2529

2630
__all__ = [
2731
"ConnectorAuthTaskState",
@@ -42,4 +46,6 @@
4246
"EncryptedStrField",
4347
"configure_workflow_encoding",
4448
"generate_two_part_id",
49+
"StreamDisconnectedError",
50+
"StreamDisconnectReason",
4551
]
Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,23 @@
1+
from __future__ import annotations
2+
3+
from typing import Literal
4+
5+
from mistralai.extra.exceptions import MistralClientException
6+
7+
StreamDisconnectReason = Literal["read_error", "stream_error", "internal_error"]
8+
9+
10+
class StreamDisconnectedError(MistralClientException):
11+
"""Raised when a workflow SSE stream is terminated by a server error frame.
12+
13+
The server ends a stream by emitting an ``event: error`` SSE frame. The SDK
14+
surfaces this as a raised exception so consumers can wrap stream iteration in
15+
``try`` / ``except`` instead of inspecting each event for ``event == "error"``.
16+
17+
Both attributes are populated from the frame's ``data`` JSON payload.
18+
"""
19+
20+
def __init__(self, *, reason: StreamDisconnectReason, error: str) -> None:
21+
self.reason: StreamDisconnectReason = reason
22+
self.error = error
23+
super().__init__("Workflow stream disconnected by server")
Lines changed: 189 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,189 @@
1+
from __future__ import annotations
2+
3+
import json
4+
import re
5+
from typing import Any, AsyncIterator, Iterator, Optional, Tuple, Union
6+
7+
import httpx
8+
from httpx._types import AsyncByteStream, SyncByteStream
9+
10+
from mistralai.client._hooks.types import AfterSuccessContext, AfterSuccessHook
11+
from mistralai.extra.workflows.errors import (
12+
StreamDisconnectReason,
13+
StreamDisconnectedError,
14+
)
15+
16+
# Operation IDs of the two SSE-backed workflow stream endpoints.
17+
STREAM_OPERATIONS = {
18+
"get_stream_events_v1_workflows_events_stream_get",
19+
"stream_v1_workflows_executions__execution_id__stream_get",
20+
}
21+
22+
_ERROR_EVENT = "error"
23+
_VALID_REASONS = ("read_error", "stream_error", "internal_error")
24+
_DEFAULT_REASON: StreamDisconnectReason = "stream_error"
25+
26+
# SSE frame boundaries (blank line), longest first so the full separator is consumed.
27+
_BOUNDARIES = [
28+
b"\r\n\r\n",
29+
b"\r\n\r",
30+
b"\r\n\n",
31+
b"\r\r\n",
32+
b"\n\r\n",
33+
b"\r\r",
34+
b"\n\r",
35+
b"\n\n",
36+
]
37+
38+
39+
def _strip_content_encoding_header(headers: httpx.Headers) -> httpx.Headers:
40+
return httpx.Headers(
41+
[(k, v) for k, v in headers.items() if k.lower() != "content-encoding"]
42+
)
43+
44+
45+
def _find_boundary(buffer: bytearray) -> Optional[Tuple[int, int]]:
46+
"""Return (index, length) of the earliest frame boundary, or None if incomplete."""
47+
best: Optional[Tuple[int, int]] = None
48+
for boundary in _BOUNDARIES:
49+
idx = buffer.find(boundary)
50+
if idx == -1:
51+
continue
52+
if (
53+
best is None
54+
or idx < best[0]
55+
or (idx == best[0] and len(boundary) > best[1])
56+
):
57+
best = (idx, len(boundary))
58+
return best
59+
60+
61+
def _parse_error_payload(data: str) -> Tuple[str, StreamDisconnectReason]:
62+
payload: dict[str, Any] = {}
63+
try:
64+
parsed = json.loads(data.strip())
65+
if isinstance(parsed, dict):
66+
payload = parsed
67+
except json.JSONDecodeError:
68+
pass
69+
error = str(payload.get("error", data.strip()))
70+
reason = payload.get("reason", _DEFAULT_REASON)
71+
if reason not in _VALID_REASONS:
72+
reason = _DEFAULT_REASON
73+
return error, reason
74+
75+
76+
def _raise_if_error_frame(block: bytes) -> None:
77+
"""Raise StreamDisconnectedError if the SSE frame is an ``event: error`` frame."""
78+
event_name: Optional[str] = None
79+
data = ""
80+
for line in re.split(r"\r?\n|\r", block.decode("utf-8", errors="replace")):
81+
if not line or line.startswith(":"):
82+
continue
83+
field, _, value = line.partition(":")
84+
if value.startswith(" "):
85+
value = value[1:]
86+
if field == "event":
87+
event_name = value
88+
elif field == "data":
89+
data += value + "\n"
90+
91+
if event_name != _ERROR_EVENT:
92+
return
93+
94+
error, reason = _parse_error_payload(data)
95+
raise StreamDisconnectedError(reason=reason, error=error)
96+
97+
98+
class _FrameScanner:
99+
"""Buffers raw SSE bytes, raising on error frames and passing others through."""
100+
101+
def __init__(self) -> None:
102+
self._buffer = bytearray()
103+
104+
def feed(self, chunk: bytes) -> Iterator[bytes]:
105+
self._buffer += chunk
106+
while True:
107+
found = _find_boundary(self._buffer)
108+
if found is None:
109+
return
110+
idx, length = found
111+
block = bytes(self._buffer[:idx])
112+
frame = bytes(self._buffer[: idx + length])
113+
del self._buffer[: idx + length]
114+
_raise_if_error_frame(block)
115+
yield frame
116+
117+
def flush(self) -> Iterator[bytes]:
118+
if not self._buffer:
119+
return
120+
block = bytes(self._buffer)
121+
self._buffer.clear()
122+
_raise_if_error_frame(block)
123+
yield block
124+
125+
126+
class _ErrorDetectingSyncByteStream(SyncByteStream):
127+
def __init__(self, original: SyncByteStream) -> None:
128+
self._original = original
129+
self._scanner = _FrameScanner()
130+
131+
def __iter__(self) -> Iterator[bytes]:
132+
for chunk in self._original:
133+
yield from self._scanner.feed(chunk)
134+
yield from self._scanner.flush()
135+
136+
def close(self) -> None:
137+
self._original.close()
138+
139+
140+
class _ErrorDetectingAsyncByteStream(AsyncByteStream):
141+
def __init__(self, original: AsyncByteStream) -> None:
142+
self._original = original
143+
self._scanner = _FrameScanner()
144+
145+
async def __aiter__(self) -> AsyncIterator[bytes]:
146+
async for chunk in self._original:
147+
for frame in self._scanner.feed(chunk):
148+
yield frame
149+
for frame in self._scanner.flush():
150+
yield frame
151+
152+
async def aclose(self) -> None:
153+
await self._original.aclose()
154+
155+
156+
class WorkflowStreamErrorHook(AfterSuccessHook):
157+
"""Raise StreamDisconnectedError when a workflow SSE stream sends an error frame.
158+
159+
Wraps the response byte stream for the two workflow SSE operations so that an
160+
``event: error`` frame raises during iteration, terminating the consumer's
161+
``for event in stream`` loop instead of yielding the error as a normal event.
162+
"""
163+
164+
def after_success(
165+
self,
166+
hook_ctx: AfterSuccessContext,
167+
response: httpx.Response,
168+
) -> Union[httpx.Response, Exception]:
169+
if hook_ctx.operation_id not in STREAM_OPERATIONS:
170+
return response
171+
if "text/event-stream" not in response.headers.get("content-type", ""):
172+
return response
173+
174+
stream = response.stream
175+
wrapped: Union[SyncByteStream, AsyncByteStream]
176+
if isinstance(stream, AsyncByteStream):
177+
wrapped = _ErrorDetectingAsyncByteStream(stream)
178+
elif isinstance(stream, SyncByteStream):
179+
wrapped = _ErrorDetectingSyncByteStream(stream)
180+
else:
181+
return response
182+
183+
return httpx.Response(
184+
status_code=response.status_code,
185+
headers=_strip_content_encoding_header(response.headers),
186+
stream=wrapped,
187+
request=response.request,
188+
extensions=response.extensions,
189+
)

0 commit comments

Comments
 (0)