diff --git a/agent/bedrock_adapter.py b/agent/bedrock_adapter.py index 5a938ff261..2771372ae3 100644 --- a/agent/bedrock_adapter.py +++ b/agent/bedrock_adapter.py @@ -824,6 +824,65 @@ def call_converse( return normalize_converse_response(response) +# Public API kept from main (plugins may import it): the agent loop itself streams through +# chat_completion_helpers._bedrock_converse_call, which applies the same recovery ladder. +def call_converse_stream( + region: str, + model: str, + messages: List[Dict], + tools: Optional[List[Dict]] = None, + max_tokens: Optional[int] = 4096, + temperature: Optional[float] = None, + top_p: Optional[float] = None, + stop_sequences: Optional[List[str]] = None, + guardrail_config: Optional[Dict] = None, +) -> SimpleNamespace: + """Call Bedrock ConverseStream API and return an OpenAI-compatible response. + + Consumes the full stream and returns the assembled response. For true + streaming with delta callbacks, use ``iter_converse_stream()`` instead. + """ + client = _get_bedrock_runtime_client(region) + kwargs = build_converse_kwargs( + model=model, + messages=messages, + tools=tools, + max_tokens=max_tokens, + temperature=temperature, + top_p=top_p, + stop_sequences=stop_sequences, + guardrail_config=guardrail_config, + ) + + try: + response = client.converse_stream(**kwargs) + except Exception as exc: + retry_kwargs = recover_from_cache_point_rejection(exc, kwargs) + if retry_kwargs is not None: + return normalize_converse_stream_events( + client.converse_stream(**retry_kwargs) + ) + if is_streaming_access_denied_error(exc): + # IAM allows bedrock:InvokeModel but not + # InvokeModelWithResponseStream — permanent for this session. + # Fall back to the non-streaming converse() path. + logger.info( + "bedrock: converse_stream denied by IAM on (region=%s, model=%s) — " + "falling back to non-streaming converse().", + region, model, + ) + return normalize_converse_response(client.converse(**kwargs)) + if is_stale_connection_error(exc): + logger.warning( + "bedrock: stale-connection error on converse_stream(region=%s, " + "model=%s): %s — evicting cached client so the next call reconnects.", + region, model, type(exc).__name__, + ) + invalidate_runtime_client(region) + raise + return normalize_converse_stream_events(response) + + # --- Model discovery --- _discovery_cache: Dict[str, Any] = {} @@ -913,6 +972,62 @@ def _extract_provider_from_arn(arn: str) -> str: return match.group(1) if match else "" +# --------------------------------------------------------------------------- +# Error classification — Bedrock-specific exceptions +# --------------------------------------------------------------------------- +# Mirrors OpenClaw's classifyFailoverReason() and matchesContextOverflowError() +# in extensions/amazon-bedrock/register.sync.runtime.ts. + +# Patterns that indicate the input context exceeded the model's token limit. +# Used by run_agent.py to trigger context compression instead of retrying. +CONTEXT_OVERFLOW_PATTERNS = [ + re.compile(r"ValidationException.*(?:input is too long|max input token|input token.*exceed)", re.IGNORECASE), + re.compile(r"ValidationException.*(?:exceeds? the (?:maximum|max) (?:number of )?(?:input )?tokens)", re.IGNORECASE), + re.compile(r"ModelStreamErrorException.*(?:Input is too long|too many input tokens)", re.IGNORECASE), +] + +# Patterns for throttling / rate limit errors — should trigger backoff + retry. +THROTTLE_PATTERNS = [ + re.compile(r"ThrottlingException", re.IGNORECASE), + re.compile(r"Too many concurrent requests", re.IGNORECASE), + re.compile(r"ServiceQuotaExceededException", re.IGNORECASE), +] + +# Patterns for transient overload — model is temporarily unavailable. +OVERLOAD_PATTERNS = [ + re.compile(r"ModelNotReadyException", re.IGNORECASE), + re.compile(r"ModelTimeoutException", re.IGNORECASE), + re.compile(r"InternalServerException", re.IGNORECASE), +] + + +def is_context_overflow_error(error_message: str) -> bool: + """Return True if the error indicates the input context was too large. + + When this returns True, the agent should compress context and retry + rather than treating it as a fatal error. + """ + return any(p.search(error_message) for p in CONTEXT_OVERFLOW_PATTERNS) + + +def classify_bedrock_error(error_message: str) -> str: + """Classify a Bedrock error for retry/failover decisions. + + Returns: + - ``"context_overflow"`` — input too long, compress and retry + - ``"rate_limit"`` — throttled, backoff and retry + - ``"overloaded"`` — model temporarily unavailable, retry with delay + - ``"unknown"`` — unclassified error + """ + if is_context_overflow_error(error_message): + return "context_overflow" + if any(p.search(error_message) for p in THROTTLE_PATTERNS): + return "rate_limit" + if any(p.search(error_message) for p in OVERLOAD_PATTERNS): + return "overloaded" + return "unknown" + + # --- Bedrock model context lengths --- # Static fallback when the live probe is unavailable (agent/model_metadata.py). Keys match by longest # substring, so versioned entries win over the generic "anthropic.claude-opus-4". diff --git a/tests/agent/test_bedrock_adapter.py b/tests/agent/test_bedrock_adapter.py index 7d5653edc3..1546369787 100644 --- a/tests/agent/test_bedrock_adapter.py +++ b/tests/agent/test_bedrock_adapter.py @@ -574,6 +574,27 @@ class TestBuildConverseKwargs: ) assert "inferenceConfig" not in kwargs + def test_call_converse_stream_omits_cap_for_none(self): + """The streaming entry point funnels through the same builder — pin + that max_tokens=None omits the cap there too.""" + from unittest.mock import MagicMock, patch as mock_patch + from agent.bedrock_adapter import call_converse_stream + boto3_client = MagicMock() + boto3_client.converse_stream.return_value = {"stream": []} + with mock_patch( + "agent.bedrock_adapter._get_bedrock_runtime_client", + return_value=boto3_client, + ): + call_converse_stream( + region="us-east-1", + model="test-model", + messages=[{"role": "user", "content": "Hi"}], + max_tokens=None, + temperature=0.2, + ) + wire_kwargs = boto3_client.converse_stream.call_args.kwargs + assert "maxTokens" not in wire_kwargs.get("inferenceConfig", {}) + def test_cache_point_added_for_supported_model(self): """Claude and Nova on the Converse path get cachePoint markers on system, tools, and the message before the newest turn.""" @@ -1004,6 +1025,20 @@ class TestGuardrailConfig: assert "guardrailConfig" not in kwargs +# --------------------------------------------------------------------------- +# Error classification +# --------------------------------------------------------------------------- + +class TestBedrockErrorClassification: + """Test Bedrock-specific error classification.""" + + def test_context_overflow_validation_exception(self): + from agent.bedrock_adapter import classify_bedrock_error + assert classify_bedrock_error( + "ValidationException: input is too long for model" + ) == "context_overflow" + + class TestBedrockContextLength: """Test Bedrock model context length lookup.""" @@ -1227,11 +1262,36 @@ class TestIsStaleConnectionError: class TestCallConverseInvalidatesOnStaleError: - """call_converse evicts the cached client when the + """call_converse / call_converse_stream evict the cached client when the boto3 call raises a stale-connection error — so the next invocation reconnects instead of reusing the dead socket.""" + def test_converse_stream_evicts_client_on_stale_error(self): + pytest.importorskip("botocore.exceptions", reason="botocore (with working exceptions module) required") + from agent.bedrock_adapter import ( + _bedrock_runtime_client_cache, + call_converse_stream, + reset_client_cache, + ) + from botocore.exceptions import ConnectionClosedError + + reset_client_cache() + dead_client = MagicMock() + dead_client.converse_stream.side_effect = ConnectionClosedError( + endpoint_url="https://bedrock.example", + ) + _bedrock_runtime_client_cache["us-east-1"] = dead_client + + with pytest.raises(ConnectionClosedError): + call_converse_stream( + region="us-east-1", + model="anthropic.claude-3-sonnet-20240229-v1:0", + messages=[{"role": "user", "content": "hi"}], + ) + + assert "us-east-1" not in _bedrock_runtime_client_cache + def test_converse_does_not_evict_on_non_stale_error(self): """Non-stale errors (e.g. ValidationException) leave the client cache alone.""" pytest.importorskip("botocore.exceptions", reason="botocore (with working exceptions module) required") @@ -1301,6 +1361,114 @@ class TestStreamingAccessDeniedDetection: ) is False +class TestCallConverseStreamIamFallback: + """call_converse_stream() falls back to converse() when IAM denies the + streaming action — InvokeModel-only policies keep working.""" + + def test_falls_back_to_converse_on_streaming_denial(self): + pytest.importorskip("botocore.exceptions", reason="botocore (with working exceptions module) required") + from agent.bedrock_adapter import ( + _bedrock_runtime_client_cache, + call_converse_stream, + reset_client_cache, + ) + from botocore.exceptions import ClientError + + reset_client_cache() + client = MagicMock() + client.converse_stream.side_effect = ClientError( + error_response={ + "Error": { + "Code": "AccessDeniedException", + "Message": ( + "User is not authorized to perform: " + "bedrock:InvokeModelWithResponseStream" + ), + } + }, + operation_name="ConverseStream", + ) + client.converse.return_value = { + "output": {"message": {"role": "assistant", "content": [{"text": "hi"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 1, "outputTokens": 1, "totalTokens": 2}, + } + _bedrock_runtime_client_cache["us-east-1"] = client + + result = call_converse_stream( + region="us-east-1", + model="anthropic.claude-3-sonnet-20240229-v1:0", + messages=[{"role": "user", "content": "hi"}], + ) + + client.converse.assert_called_once() + assert result.choices[0].message.content == "hi" + # Not a stale connection — client stays cached. + assert _bedrock_runtime_client_cache.get("us-east-1") is client + + +class TestAgentBedrockStreamRecovery: + """The agent loop streams through ``chat_completion_helpers._bedrock_converse_call`` + (not ``call_converse_stream``); pin the same recovery ladder on that live path: + IAM streaming denial → ``_BedrockStream._fall_back_to_converse`` (non-streaming + converse, streaming disabled for the session, client kept), stale connection → + cached client evicted so the outer retry reconnects.""" + + _KW = {"__bedrock_region__": "us-east-1", "modelId": "anthropic.claude-3-sonnet-20240229-v1:0", + "messages": [{"role": "user", "content": [{"text": "hi"}]}]} + + def test_streaming_denial_falls_back_to_converse_via_bedrock_stream(self): + pytest.importorskip("botocore.exceptions", reason="botocore (with working exceptions module) required") + from types import SimpleNamespace + from agent.bedrock_adapter import _bedrock_runtime_client_cache, reset_client_cache + from agent.chat_completion_helpers import _BedrockStream + from botocore.exceptions import ClientError + + reset_client_cache() + client = MagicMock() + client.converse_stream.side_effect = ClientError( + error_response={"Error": {"Code": "AccessDeniedException", "Message": ( + "User is not authorized to perform: bedrock:InvokeModelWithResponseStream")}}, + operation_name="ConverseStream", + ) + client.converse.return_value = { + "output": {"message": {"role": "assistant", "content": [{"text": "hi"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 1, "outputTokens": 1, "totalTokens": 2}, + } + _bedrock_runtime_client_cache["us-east-1"] = client + agent = SimpleNamespace(_disable_streaming=False, _safe_print=MagicMock(), model="m", provider="bedrock") + stream = _BedrockStream(agent, dict(self._KW), on_first_delta=None) + + result = stream._open_stream(dict(self._KW)) + + client.converse.assert_called_once() + assert "__bedrock_region__" not in client.converse.call_args.kwargs + assert result.choices[0].message.content == "hi" + assert agent._disable_streaming is True + assert "InvokeModelWithResponseStream" in agent._safe_print.call_args.args[0] + # Not a stale connection — client stays cached. + assert _bedrock_runtime_client_cache.get("us-east-1") is client + + def test_stale_connection_evicts_client_on_agent_stream_path(self): + pytest.importorskip("botocore.exceptions", reason="botocore (with working exceptions module) required") + from agent.bedrock_adapter import _bedrock_runtime_client_cache, reset_client_cache + from agent.chat_completion_helpers import _bedrock_converse_call + from botocore.exceptions import ConnectionClosedError + + reset_client_cache() + dead_client = MagicMock() + dead_client.converse_stream.side_effect = ConnectionClosedError(endpoint_url="https://bedrock.example") + _bedrock_runtime_client_cache["us-east-1"] = dead_client + denied = MagicMock() + + with pytest.raises(ConnectionClosedError): + _bedrock_converse_call(dict(self._KW), stream=True, on_stream_denied=denied) + + denied.assert_not_called() + assert "us-east-1" not in _bedrock_runtime_client_cache + + # --------------------------------------------------------------------------- # boto3 version check # ---------------------------------------------------------------------------