Skip to content

Commit 472daf2

Browse files
refactor(workflows): move stream error hook into _hooks and cover logs streams
- Move WorkflowStreamErrorHook to client/_hooks alongside the other hooks - Cover deployment + execution logs SSE streams (same error-frame contract) - Derive valid reasons from the StreamDisconnectReason Literal - Add tests: error frame without trailing boundary; all stream operations
1 parent a7d9e33 commit 472daf2

3 files changed

Lines changed: 42 additions & 10 deletions

File tree

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

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,8 +3,8 @@
33
from .traceparent import TraceparentInjectionHook
44
from .tracing import TracingHook
55
from .types import Hooks
6+
from .stream_error_hook import WorkflowStreamErrorHook
67
from .workflow_encoding_hook import WorkflowEncodingHook
7-
from mistralai.extra.workflows.stream_error_hook import WorkflowStreamErrorHook
88

99
# This file is only ever generated once on the first generation and then is free to be modified.
1010
# Any hooks you wish to add should be registered in the init_hooks function. Feel free to define them

src/mistralai/extra/workflows/stream_error_hook.py renamed to src/mistralai/client/_hooks/stream_error_hook.py

Lines changed: 9 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1,26 +1,27 @@
1-
from __future__ import annotations
2-
31
import json
42
import re
5-
from typing import Any, AsyncIterator, Iterator, Optional, Tuple, Union
3+
from typing import Any, AsyncIterator, Dict, Iterator, Optional, Tuple, Union, get_args
64

75
import httpx
86
from httpx._types import AsyncByteStream, SyncByteStream
97

10-
from mistralai.client._hooks.types import AfterSuccessContext, AfterSuccessHook
8+
from .types import AfterSuccessContext, AfterSuccessHook
119
from mistralai.extra.exceptions import (
1210
StreamDisconnectReason,
1311
StreamDisconnectedError,
1412
)
1513

16-
# Operation IDs of the two SSE-backed workflow stream endpoints.
14+
# Operation IDs of the SSE-backed workflow stream endpoints that can emit a
15+
# terminal ``event: error`` frame (event, execution, and logs streams).
1716
STREAM_OPERATIONS = {
1817
"get_stream_events_v1_workflows_events_stream_get",
1918
"stream_v1_workflows_executions__execution_id__stream_get",
19+
"stream_deployment_logs",
20+
"stream_workflow_execution_logs",
2021
}
2122

2223
_ERROR_EVENT = "error"
23-
_VALID_REASONS = ("read_error", "stream_error", "internal_error")
24+
_VALID_REASONS = get_args(StreamDisconnectReason)
2425
_DEFAULT_REASON: StreamDisconnectReason = "stream_error"
2526

2627
# SSE frame boundaries (blank line), longest first so the full separator is consumed.
@@ -59,7 +60,7 @@ def _find_boundary(buffer: bytearray) -> Optional[Tuple[int, int]]:
5960

6061

6162
def _parse_error_payload(data: str) -> Tuple[str, StreamDisconnectReason]:
62-
payload: dict[str, Any] = {}
63+
payload: Dict[str, Any] = {}
6364
try:
6465
parsed = json.loads(data.strip())
6566
if isinstance(parsed, dict):
@@ -156,7 +157,7 @@ async def aclose(self) -> None:
156157
class WorkflowStreamErrorHook(AfterSuccessHook):
157158
"""Raise StreamDisconnectedError when a workflow SSE stream sends an error frame.
158159
159-
Wraps the response byte stream for the two workflow SSE operations so that an
160+
Wraps the response byte stream for the workflow SSE operations so that an
160161
``event: error`` frame raises during iteration, terminating the consumer's
161162
``for event in stream`` loop instead of yielding the error as a normal event.
162163
"""

‎src/mistralai/extra/tests/test_stream_error_hook.py‎

Lines changed: 32 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,10 @@
55
from mistralai.client import Mistral
66
from mistralai.client._hooks.types import AfterSuccessContext, HookContext
77
from mistralai.extra.exceptions import StreamDisconnectedError
8-
from mistralai.extra.workflows.stream_error_hook import WorkflowStreamErrorHook
8+
from mistralai.client._hooks.stream_error_hook import (
9+
STREAM_OPERATIONS,
10+
WorkflowStreamErrorHook,
11+
)
912

1013
STREAM_OPERATION_ID = "get_stream_events_v1_workflows_events_stream_get"
1114
NON_STREAM_OPERATION_ID = "chat_completion_v1_chat_completions_post"
@@ -119,6 +122,21 @@ def test_hook_detects_error_frame_split_across_chunks():
119122
assert exc_info.value.error == "x"
120123

121124

125+
def test_hook_raises_on_error_frame_without_trailing_boundary():
126+
frame = b'event: error\ndata: {"error": "boom", "reason": "read_error"}'
127+
response = _sse_response(_SyncSource([frame]))
128+
result = WorkflowStreamErrorHook().after_success(
129+
_hook_ctx(STREAM_OPERATION_ID), response
130+
)
131+
assert isinstance(result, httpx.Response)
132+
133+
with pytest.raises(StreamDisconnectedError) as exc_info:
134+
list(result.iter_bytes())
135+
136+
assert exc_info.value.reason == "read_error"
137+
assert exc_info.value.error == "boom"
138+
139+
122140
def test_hook_defaults_reason_when_missing_or_invalid():
123141
frame = b'event: error\ndata: {"error": "no reason given"}\n\n'
124142
response = _sse_response(_SyncSource([frame]))
@@ -145,6 +163,19 @@ def test_hook_passes_normal_stream_through_without_raising():
145163
assert body.count(b"workflow.event") == 2
146164

147165

166+
@pytest.mark.parametrize("operation_id", sorted(STREAM_OPERATIONS))
167+
def test_hook_raises_for_every_workflow_stream_operation(operation_id: str):
168+
response = _sse_response(_SyncSource([ERROR_FRAME]))
169+
result = WorkflowStreamErrorHook().after_success(_hook_ctx(operation_id), response)
170+
assert isinstance(result, httpx.Response)
171+
172+
with pytest.raises(StreamDisconnectedError) as exc_info:
173+
list(result.iter_bytes())
174+
175+
assert exc_info.value.reason == "read_error"
176+
assert exc_info.value.error == "boom"
177+
178+
148179
def test_hook_ignores_non_stream_operations():
149180
source = _SyncSource([ERROR_FRAME])
150181
response = _sse_response(source)

0 commit comments

Comments
 (0)