diff --git a/api/streaming.py b/api/streaming.py index ff07a72d2..88c3cf495 100644 --- a/api/streaming.py +++ b/api/streaming.py @@ -4440,6 +4440,96 @@ def _materialize_pending_user_turn_before_error(session) -> bool: return True +def _build_partial_message(content_text, reasoning_text, tool_calls) -> dict | None: + """Build a _partial assistant message from raw streaming buffers. + + Shared by cancel_stream() and _snapshot_and_append_partial_on_error(). + Strips thinking/reasoning markup, builds the dict, returns None when + there is nothing meaningful to preserve. + """ + import re as _re + partial_text = (content_text or '').strip() + _stripped = '' + if partial_text: + # First pass: remove complete ... blocks. + _stripped = _re.sub(r']*>.*?', + '', partial_text, + flags=_re.DOTALL | _re.IGNORECASE).strip() + # Second pass: strip trailing UNCLOSED think/thinking block (the common + # cancel/error case — user stops mid-reasoning before the close tag appears). + _stripped = _re.sub(r']*>.*', + '', _stripped, + flags=_re.DOTALL | _re.IGNORECASE).strip() + _has_reasoning = bool(reasoning_text and reasoning_text.strip()) + _has_tools = bool(tool_calls) + if not (_stripped or _has_reasoning or _has_tools): + return None + _msg: dict = { + 'role': 'assistant', + 'content': _stripped, # may be empty for reasoning/tool-only turns + '_partial': True, + 'timestamp': int(time.time()), + } + if _has_reasoning: + _msg['reasoning'] = reasoning_text.strip() + if _has_tools: + _msg['_partial_tool_calls'] = list(tool_calls) + return _msg + + +def _snapshot_and_append_partial_on_error(session, stream_id) -> dict | None: + """Snapshot streaming buffers under STREAMS_LOCK and append a _partial message. + + Uses _build_partial_message() for the shared thinking-strip + dict-build logic. + """ + from api import config as _live_config + + streams_lock = STREAMS_LOCK + partial_texts = STREAM_PARTIAL_TEXT + reasoning_texts = STREAM_REASONING_TEXT + live_tool_calls = STREAM_LIVE_TOOL_CALLS + + # Defensive check for live config (similar to cancel_stream) + if getattr(_live_config, 'STREAMS_LOCK', streams_lock) is not streams_lock: + streams_lock = _live_config.STREAMS_LOCK + partial_texts = getattr(_live_config, 'STREAM_PARTIAL_TEXT', partial_texts) + reasoning_texts = getattr(_live_config, 'STREAM_REASONING_TEXT', reasoning_texts) + live_tool_calls = getattr(_live_config, 'STREAM_LIVE_TOOL_CALLS', live_tool_calls) + + _snap_partial_text = None + _snap_reasoning = None + _snap_tool_calls = None + + with streams_lock: + _snap_partial_text = partial_texts.get(stream_id, '') + if not _snap_partial_text: + _live_partials = getattr(_live_config, 'STREAM_PARTIAL_TEXT', partial_texts) + if _live_partials is not partial_texts: + _snap_partial_text = _live_partials.get(stream_id, '') + + _snap_reasoning = reasoning_texts.get(stream_id, '') + if not _snap_reasoning: + _live_reasoning = getattr(_live_config, 'STREAM_REASONING_TEXT', reasoning_texts) + if _live_reasoning is not reasoning_texts: + _snap_reasoning = _live_reasoning.get(stream_id, '') + + _snap_tool_calls = list(live_tool_calls.get(stream_id, []) or []) + if not _snap_tool_calls: + _live_tools = getattr(_live_config, 'STREAM_LIVE_TOOL_CALLS', live_tool_calls) + if _live_tools is not live_tool_calls: + _snap_tool_calls = list(_live_tools.get(stream_id, []) or []) + + _partial_msg = _build_partial_message(_snap_partial_text, _snap_reasoning, _snap_tool_calls) + if _partial_msg is None: + return None + if not isinstance(session.messages, list): + session.messages = [] + if not _partial_marker_already_present(session.messages, _partial_msg): + session.messages.append(_partial_msg) + return _partial_msg + return None + + def _last_resort_sync_from_core(session, stream_id, agent_lock): """Final-exit guard: if the stream exits with pending_user_message still set, sync messages from the core transcript or add an error marker. @@ -6787,16 +6877,15 @@ def _run_agent_streaming( _err_type, _err_hint, ) - # Clear stream/pending state so the session does not appear - # "agent_running" on reload after a silent failure. - # Persist the error so it survives page reload. - # _error=True ensures _sanitize_messages_for_api excludes it from - # subsequent API calls so the LLM never sees its own error as prior context. _materialize_pending_user_turn_before_error(s) s.active_stream_id = None s.pending_user_message = None s.pending_attachments = [] s.pending_started_at = None + try: + _snapshot_and_append_partial_on_error(s, stream_id) + except Exception: + logger.debug("Failed to snapshot partials on error for %s", stream_id, exc_info=True) _error_message = { 'role': 'assistant', 'content': f'**{_err_label}:** {_error_payload.get("message") or _err_label}\n\n*{_err_hint}*', @@ -7741,11 +7830,16 @@ def _run_agent_streaming( getattr(s, 'active_stream_id', None), ) return + _materialize_pending_user_turn_before_error(s) s.active_stream_id = None s.pending_user_message = None s.pending_attachments = [] s.pending_started_at = None + try: + _snapshot_and_append_partial_on_error(s, stream_id) + except Exception: + logger.debug("Failed to snapshot partials on error for %s", stream_id, exc_info=True) _error_message = { 'role': 'assistant', 'content': f'**{_exc_label}:** {_error_payload.get("message") or err_str}' + (f'\n\n*{_exc_hint}*' if _exc_hint else ''), @@ -8233,26 +8327,17 @@ def cancel_stream(stream_id: str) -> bool: # accumulated in thread-local variables but invisible to the cancel path. # This prevents paid-token data loss when cancelling mid-reasoning or # mid-tool-execution. - partial_text = _cancel_partial_text.strip() if _cancel_partial_text else '' - _stripped = '' - if partial_text: - import re as _re - # Strip thinking/reasoning markup from partial content before saving. - # First pass: remove complete ... blocks. - _stripped = _re.sub(r']*>.*?', - '', partial_text, - flags=_re.DOTALL | _re.IGNORECASE).strip() - # Second pass: strip trailing UNCLOSED think/thinking block (the common - # cancel case — user stops mid-reasoning before the close tag appears). - _stripped = _re.sub(r']*>.*', - '', _stripped, - flags=_re.DOTALL | _re.IGNORECASE).strip() - # Determine whether there is anything to preserve beyond just the - # cancel marker. Content text, reasoning trace, or tool calls all - # count (#1361 §C — previously only _stripped was checked, so a - # reasoning-only or tool-only stream produced NO partial message). - _has_reasoning = bool(_cancel_reasoning and _cancel_reasoning.strip()) - _has_tools = bool(_cancel_tool_calls) + # NOTE on _partial_tool_calls: the captured entries use the WebUI + # internal shape {name, args, done, duration, is_error} — they do + # NOT carry the OpenAI/Anthropic API id + function: {name, arguments} + # envelope. Storing under 'tool_calls' would cause + # _sanitize_messages_for_api to forward them to the next-turn LLM + # call and strict providers would 400 on the malformed entries. + # The underscore-prefixed key is not in the whitelist, so sanitize + # strips it. The UI reads it via static/messages.js. (v0.50.251.) + _partial_msg = _build_partial_message( + _cancel_partial_text, _cancel_reasoning, _cancel_tool_calls, + ) _cancel_marker_exists = _session_has_cancel_marker(_cs) _cancel_marker_idx = len(_cs.messages) if _cancel_marker_exists: @@ -8264,33 +8349,7 @@ def cancel_stream(stream_id: str) -> bool: if any(pattern in _content for pattern in _CANCEL_MARKER_PATTERNS): _cancel_marker_idx = _idx break - if _stripped or _has_reasoning or _has_tools: - _partial_msg: dict = { - 'role': 'assistant', - 'content': _stripped, # may be empty for reasoning/tool-only turns - '_partial': True, - 'timestamp': int(time.time()), - } - if _has_reasoning: - _partial_msg['reasoning'] = _cancel_reasoning.strip() - if _has_tools: - # NOTE: store under the private '_partial_tool_calls' key - # (NOT 'tool_calls'). The captured entries use the WebUI - # internal shape {name, args, done, duration, is_error} - # — they do NOT carry the OpenAI/Anthropic API id + - # function: {name, arguments} envelope. If we put them - # under 'tool_calls', `_sanitize_messages_for_api` - # (which whitelists 'tool_calls' via _API_SAFE_MSG_KEYS) - # would forward them to the next-turn LLM call and - # strict providers (OpenAI, Anthropic, Z.AI/GLM) would - # 400 on the malformed entries — turning a "data lost - # on cancel" bug into a "next message returns 400" - # bug, which is worse. The underscore-prefixed key is - # not in the whitelist, so sanitize strips it. The UI - # reads it via static/messages.js and renders it - # alongside the regular tool_calls path. - # (Opus pre-release review pass 2 of v0.50.251.) - _partial_msg['_partial_tool_calls'] = list(_cancel_tool_calls) + if _partial_msg is not None: # Deduplicate against the full partial payload, not just # non-empty content. Tool-only/reasoning-only partials have # empty content, so a content-gated check can append the same diff --git a/tests/test_issue3929_error_preserves_partial.py b/tests/test_issue3929_error_preserves_partial.py new file mode 100644 index 000000000..6584cf24b --- /dev/null +++ b/tests/test_issue3929_error_preserves_partial.py @@ -0,0 +1,247 @@ +"""Tests for Issue #3929: Ensure turn-level error handling preserves partial work in the WebUI session sidecar.""" + +from __future__ import annotations + +import queue +import sys +import types +import pytest +from unittest import mock + +import api.config as config +import api.models as models +import api.streaming as streaming +from api.models import Session + + +@pytest.fixture(autouse=True) +def _isolate_session_dir(tmp_path, monkeypatch): + """Redirect SESSION_DIR / SESSION_INDEX_FILE to an isolated temp dir.""" + session_dir = tmp_path / "sessions" + session_dir.mkdir() + index_file = session_dir / "_index.json" + monkeypatch.setattr(models, "SESSION_DIR", session_dir) + monkeypatch.setattr(models, "SESSION_INDEX_FILE", index_file) + models.SESSIONS.clear() + yield + models.SESSIONS.clear() + + +@pytest.fixture(autouse=True) +def _isolate_stream_state(): + """Clear all shared streaming dicts before/after each test.""" + config.STREAMS.clear() + config.CANCEL_FLAGS.clear() + config.AGENT_INSTANCES.clear() + config.STREAM_PARTIAL_TEXT.clear() + if hasattr(config, 'STREAM_REASONING_TEXT'): + config.STREAM_REASONING_TEXT.clear() + if hasattr(config, 'STREAM_LIVE_TOOL_CALLS'): + config.STREAM_LIVE_TOOL_CALLS.clear() + yield + config.STREAMS.clear() + config.CANCEL_FLAGS.clear() + config.AGENT_INSTANCES.clear() + config.STREAM_PARTIAL_TEXT.clear() + if hasattr(config, 'STREAM_REASONING_TEXT'): + config.STREAM_REASONING_TEXT.clear() + if hasattr(config, 'STREAM_LIVE_TOOL_CALLS'): + config.STREAM_LIVE_TOOL_CALLS.clear() + + +@pytest.fixture(autouse=True) +def _isolate_agent_locks(): + config.SESSION_AGENT_LOCKS.clear() + yield + config.SESSION_AGENT_LOCKS.clear() + + +@pytest.fixture(autouse=True) +def _mock_hermes_modules(monkeypatch): + """Inject mock hermes modules to prevent side-effects during tests.""" + fake_runtime_module = types.ModuleType("hermes_cli.runtime_provider") + fake_runtime_provider_fn = lambda requested=None: { + "provider": requested or "test-provider", + "api_key": "synthetic-key", + "base_url": None, + } + fake_runtime_module.resolve_runtime_provider = fake_runtime_provider_fn + fake_hermes_cli = types.ModuleType("hermes_cli") + fake_hermes_cli.runtime_provider = fake_runtime_module + fake_hermes_state = types.ModuleType("hermes_state") + fake_hermes_state.SessionDB = mock.Mock(return_value=None) + + _injected = { + "hermes_cli": fake_hermes_cli, + "hermes_cli.runtime_provider": fake_runtime_module, + "hermes_state": fake_hermes_state, + } + _MISSING = object() + _saved = {k: sys.modules.get(k, _MISSING) for k in _injected} + sys.modules.update(_injected) + yield + for k, prev in _saved.items(): + if prev is _MISSING: + sys.modules.pop(k, None) + else: + sys.modules[k] = prev + + +class MockAgent: + def __init__(self, **kwargs): + self.session_id = kwargs.get("session_id") + self.stream_delta_callback = kwargs.get("stream_delta_callback") + self.reasoning_callback = kwargs.get("reasoning_callback") + self.tool_progress_callback = kwargs.get("tool_progress_callback") + self.session_prompt_tokens = 0 + self.session_completion_tokens = 0 + self.session_estimated_cost_usd = 0.0 + self.context_compressor = None + self._last_error = None + self.ephemeral_system_prompt = None + + def run_conversation(self, **kwargs): + pass + + def interrupt(self, _message): + pass + + +def test_silent_failure_preserves_partials(tmp_path): + """Test that a silent failure (agent returns no assistant reply) preserves partial streamed text.""" + fake_session = Session(session_id="test_sess_silent", title="Test Session") + fake_session.pending_user_message = "What is python?" + fake_session.active_stream_id = "test_stream_silent" + fake_session.save() + models.SESSIONS["test_sess_silent"] = fake_session + + class SilentFailureAgent(MockAgent): + def run_conversation(self, **kwargs): + if self.stream_delta_callback: + self.stream_delta_callback("Python is a programming language.") + # Return history without new assistant message, plus error status (causes silent failure) + return { + "status": "error", + "error": "Silent failure details", + "messages": kwargs.get("conversation_history") or [] + } + + fake_queue = queue.Queue() + streaming.STREAMS["test_stream_silent"] = fake_queue + config.STREAM_PARTIAL_TEXT["test_stream_silent"] = "" + + with mock.patch.object(streaming, "get_session", return_value=fake_session), \ + mock.patch.object(streaming, "_get_ai_agent", return_value=SilentFailureAgent), \ + mock.patch.object(streaming, "resolve_model_provider", return_value=("test-model", "test-provider", None)), \ + mock.patch("api.config.get_config", return_value={}), \ + mock.patch("api.config._resolve_cli_toolsets", return_value=[]): + + streaming._run_agent_streaming( + session_id=fake_session.session_id, + msg_text="What is python?", + model="test-model", + workspace=str(tmp_path), + stream_id="test_stream_silent", + ) + + # Reload session from disk/cache and verify + saved = Session.load("test_sess_silent") + assert saved is not None + + # We should have: + # 1. The user turn (materialized by _materialize_pending_user_turn_before_error) + # 2. The _partial assistant message containing the streamed text + # 3. The error block + assert len(saved.messages) >= 3, f"Expected at least 3 messages. Got: {saved.messages}" + + partial_msg = next((m for m in saved.messages if m.get("_partial")), None) + assert partial_msg is not None, "Expected a partial message" + assert partial_msg["role"] == "assistant" + assert partial_msg["content"] == "Python is a programming language." + + err_msg = saved.messages[-1] + assert err_msg.get("_error") is True + + +def test_exception_preserves_partials(tmp_path): + """Test that an unhandled exception preserves partial streamed text.""" + fake_session = Session(session_id="test_sess_exc", title="Test Session") + fake_session.pending_user_message = "Exception test" + fake_session.active_stream_id = "test_stream_exc" + fake_session.save() + models.SESSIONS["test_sess_exc"] = fake_session + + class ExceptionAgent(MockAgent): + def run_conversation(self, **kwargs): + if self.stream_delta_callback: + self.stream_delta_callback("Stream before crash.") + raise RuntimeError("Fake provider crash!") + + fake_queue = queue.Queue() + streaming.STREAMS["test_stream_exc"] = fake_queue + config.STREAM_PARTIAL_TEXT["test_stream_exc"] = "" + + with mock.patch.object(streaming, "get_session", return_value=fake_session), \ + mock.patch.object(streaming, "_get_ai_agent", return_value=ExceptionAgent), \ + mock.patch.object(streaming, "resolve_model_provider", return_value=("test-model", "test-provider", None)), \ + mock.patch("api.config.get_config", return_value={}), \ + mock.patch("api.config._resolve_cli_toolsets", return_value=[]): + + streaming._run_agent_streaming( + session_id=fake_session.session_id, + msg_text="Exception test", + model="test-model", + workspace=str(tmp_path), + stream_id="test_stream_exc", + ) + + saved = Session.load("test_sess_exc") + assert saved is not None + + partial_msg = next((m for m in saved.messages if m.get("_partial")), None) + assert partial_msg is not None, "Expected a partial message before error" + assert partial_msg["content"] == "Stream before crash." + + err_msg = saved.messages[-1] + assert err_msg.get("_error") is True + assert "Fake provider crash!" in err_msg.get("content", "") + + +def test_empty_partials_do_not_create_spurious_messages(tmp_path): + """Test that if no content has been streamed, no _partial message is created on error.""" + fake_session = Session(session_id="test_sess_empty", title="Test Session") + fake_session.pending_user_message = "Empty test" + fake_session.active_stream_id = "test_stream_empty" + fake_session.save() + models.SESSIONS["test_sess_empty"] = fake_session + + class ExceptionAgent(MockAgent): + def run_conversation(self, **kwargs): + raise RuntimeError("Fake provider crash immediately!") + + fake_queue = queue.Queue() + streaming.STREAMS["test_stream_empty"] = fake_queue + config.STREAM_PARTIAL_TEXT["test_stream_empty"] = "" + + with mock.patch.object(streaming, "get_session", return_value=fake_session), \ + mock.patch.object(streaming, "_get_ai_agent", return_value=ExceptionAgent), \ + mock.patch.object(streaming, "resolve_model_provider", return_value=("test-model", "test-provider", None)), \ + mock.patch("api.config.get_config", return_value={}), \ + mock.patch("api.config._resolve_cli_toolsets", return_value=[]): + + streaming._run_agent_streaming( + session_id=fake_session.session_id, + msg_text="Empty test", + model="test-model", + workspace=str(tmp_path), + stream_id="test_stream_empty", + ) + + saved = Session.load("test_sess_empty") + assert saved is not None + + # We should have user message and error message, but NO _partial message + assert not any(m.get("_partial") for m in saved.messages) + + err_msg = saved.messages[-1] + assert err_msg.get("_error") is True