diff --git a/src/uipath_langchain/agent/tools/base_uipath_structured_tool.py b/src/uipath_langchain/agent/tools/base_uipath_structured_tool.py index 247ccfcfe..00002671f 100644 --- a/src/uipath_langchain/agent/tools/base_uipath_structured_tool.py +++ b/src/uipath_langchain/agent/tools/base_uipath_structured_tool.py @@ -138,7 +138,7 @@ def _parse_input( parsed[alias] = getattr(result, python_name) return parsed - def _invalid_input_error(self, error: ValidationError) -> AgentRuntimeError: + def _invalid_input_error(self, error: ValidationError) -> Exception: return AgentRuntimeError( code=AgentRuntimeErrorCode.INVALID_INPUT_ARGUMENT, title=f"Invalid input for tool '{self.name}'", diff --git a/src/uipath_langchain/agent/tools/internal_tools/http_request_tool.py b/src/uipath_langchain/agent/tools/internal_tools/http_request_tool.py index 303006d61..3e3143c0e 100644 --- a/src/uipath_langchain/agent/tools/internal_tools/http_request_tool.py +++ b/src/uipath_langchain/agent/tools/internal_tools/http_request_tool.py @@ -10,6 +10,12 @@ Requests that resolve to private, loopback, link-local, or cloud-metadata addresses are rejected to guard against SSRF; the check runs on every request, including redirect hops. + +Mistakes the model can correct (arguments that fail the input schema, a bad +url/method/timeout, a blocked or unreachable host, a timeout) are raised as +``ToolException`` and reported back to the model as an error tool message, so +they cost a turn instead of faulting the run. Header and param pairs accept +their ``name``/``value`` keys in any case. """ import asyncio @@ -23,17 +29,12 @@ import httpx from langchain_core.language_models import BaseChatModel -from langchain_core.tools import StructuredTool -from pydantic import BaseModel +from langchain_core.tools import StructuredTool, ToolException +from pydantic import BaseModel, ValidationError from uipath._utils._ssl_context import get_httpx_client_kwargs from uipath.agent.models.agent import AgentInternalToolResourceConfig from uipath.eval.mocks import mockable -from uipath.runtime.errors import UiPathErrorCategory -from uipath_langchain.agent.exceptions import ( - AgentRuntimeError, - AgentRuntimeErrorCode, -) from uipath_langchain.agent.react.jsonschema_pydantic_converter import ( create_model, create_output_model, @@ -73,6 +74,10 @@ # whether the caller supplied a scheme at all. _SCHEME_RE = re.compile(r"^[a-zA-Z][a-zA-Z0-9+.\-]*://") +# Arguments modeled as lists of ``{name, value}`` pairs, and the keys of a pair. +_PAIR_ARGUMENTS = ("headers", "params") +_PAIR_KEYS = ("name", "value") + def _normalize_url(url: str) -> str: """Default to ``https`` when the caller omits a scheme. @@ -111,6 +116,32 @@ def _is_json_object_body(text: str) -> bool: return isinstance(parsed, (dict, list)) +def _normalize_pair_item(item: Any) -> Any: + """Lowercase the ``name``/``value`` keys of one pair, leaving its values as-is. + + Models sometimes write ``{"Name": ..., "Value": ...}``; the schema keys are + lowercase and pydantic matches them case-sensitively. + """ + if not isinstance(item, dict): + return item + normalized: dict[Any, Any] = {} + for key, value in item.items(): + if isinstance(key, str) and key.lower() in _PAIR_KEYS: + key = key.lower() + normalized[key] = value + return normalized + + +def _normalize_pair_keys(tool_input: dict[str, Any]) -> dict[str, Any]: + """Copy the tool input with the headers/params pair keys lowercased.""" + normalized = dict(tool_input) + for argument in _PAIR_ARGUMENTS: + pairs = normalized.get(argument) + if isinstance(pairs, list): + normalized[argument] = [_normalize_pair_item(item) for item in pairs] + return normalized + + def _pairs_to_dict(pairs: list[Any]) -> dict[str, Any]: """Fold a list of ``{name, value}`` items into a dict. @@ -124,6 +155,7 @@ def _pairs_to_dict(pairs: list[Any]) -> dict[str, Any]: for item in pairs: if isinstance(item, BaseModel): item = item.model_dump() + item = _normalize_pair_item(item) if isinstance(item, dict) and item.get("name") is not None: result[str(item["name"])] = item.get("value") return result @@ -170,13 +202,8 @@ def _is_blocked_ip(ip: ipaddress.IPv4Address | ipaddress.IPv6Address) -> bool: ) -def _blocked_host_error(detail: str) -> AgentRuntimeError: - return AgentRuntimeError( - code=AgentRuntimeErrorCode.HTTP_ERROR, - title="Request was blocked", - detail=detail, - category=UiPathErrorCategory.USER, - ) +def _blocked_host_error(detail: str) -> ToolException: + return ToolException(f"Request was blocked: {detail}") async def _assert_public_url(url: str) -> None: @@ -208,11 +235,8 @@ async def _assert_public_url(url: str) -> None: host, parsed.port, type=socket.SOCK_STREAM ) except socket.gaierror as e: - raise AgentRuntimeError( - code=AgentRuntimeErrorCode.HTTP_ERROR, - title="Host could not be resolved", - detail=f"Could not resolve host {host!r}: {e}", - category=UiPathErrorCategory.USER, + raise ToolException( + f"Host could not be resolved: Could not resolve host {host!r}: {e}" ) from e for info in infos: @@ -244,18 +268,10 @@ def _validate_url(kwargs: dict[str, Any]) -> str: """Return the required, https-normalized url; raise on missing/non-string.""" url = kwargs.get("url") if not url: - raise AgentRuntimeError( - code=AgentRuntimeErrorCode.INVALID_INPUT_ARGUMENT, - title="Missing required argument", - detail="Argument 'url' is required.", - category=UiPathErrorCategory.USER, - ) + raise ToolException("Missing required argument: Argument 'url' is required.") if not isinstance(url, str): - raise AgentRuntimeError( - code=AgentRuntimeErrorCode.INVALID_INPUT_ARGUMENT, - title="Invalid url", - detail=f"Argument 'url' must be a string; got {type(url).__name__}.", - category=UiPathErrorCategory.USER, + raise ToolException( + f"Invalid url: Argument 'url' must be a string; got {type(url).__name__}." ) return _normalize_url(url) @@ -264,14 +280,9 @@ def _validate_method(kwargs: dict[str, Any]) -> str: """Return the upper-cased method (default GET); raise on unsupported.""" method = (kwargs.get("method") or "GET").upper() if method not in HTTP_REQUEST_METHODS: - raise AgentRuntimeError( - code=AgentRuntimeErrorCode.INVALID_INPUT_ARGUMENT, - title="Unsupported HTTP method", - detail=( - f"Unsupported HTTP method {method!r}; expected one of " - f"{', '.join(HTTP_REQUEST_METHODS)}." - ), - category=UiPathErrorCategory.USER, + raise ToolException( + f"Unsupported HTTP method {method!r}; expected one of " + f"{', '.join(HTTP_REQUEST_METHODS)}." ) return method @@ -286,14 +297,9 @@ def _validate_timeout(kwargs: dict[str, Any]) -> float: or not isinstance(timeout, (int, float)) or timeout <= 0 ): - raise AgentRuntimeError( - code=AgentRuntimeErrorCode.INVALID_INPUT_ARGUMENT, - title="Invalid timeout", - detail=( - "Argument 'timeout' must be a positive number of seconds; " - f"got {timeout!r}." - ), - category=UiPathErrorCategory.USER, + raise ToolException( + "Invalid timeout: Argument 'timeout' must be a positive number of " + f"seconds; got {timeout!r}." ) return timeout @@ -331,7 +337,7 @@ def _build_request_parameters(kwargs: dict[str, Any]) -> _HttpRequestParameters: """Validate and normalize the tool's input arguments into a request spec. Raises: - AgentRuntimeError: If any argument is missing or invalid (USER category). + ToolException: If any argument is missing or invalid. """ url = _validate_url(kwargs) method = _validate_method(kwargs) @@ -350,6 +356,28 @@ def _build_request_parameters(kwargs: dict[str, Any]) -> _HttpRequestParameters: ) +class _HttpRequestTool(StructuredToolWithArgumentProperties): + """HTTP request tool that tolerates pair-key casing and reports bad input. + + Arguments that still fail the input schema after the pair keys are + normalized come back to the model as an error tool message rather than + faulting the run (requires ``handle_tool_error``). + """ + + def _parse_input( + self, tool_input: str | dict[str, Any], tool_call_id: str | None + ) -> str | dict[str, Any]: + if isinstance(tool_input, dict): + tool_input = _normalize_pair_keys(tool_input) + return super()._parse_input(tool_input, tool_call_id) + + def _invalid_input_error(self, error: ValidationError) -> Exception: + return ToolException( + f"Invalid input for tool '{self.name}'. Fix the arguments so they " + f"match the tool input schema and call the tool again.\n\n{error}" + ) + + def create_http_request_tool( resource: AgentInternalToolResourceConfig, llm: BaseChatModel ) -> StructuredTool: @@ -392,24 +420,17 @@ async def http_request_tool_fn(**kwargs: Any) -> dict[str, Any]: timeout=request_parameters.timeout, **request_parameters.request_kwargs, ) - except AgentRuntimeError: + except ToolException: raise except httpx.TimeoutException as e: - raise AgentRuntimeError( - code=AgentRuntimeErrorCode.HTTP_ERROR, - title="HTTP request timed out", - detail=( - f"Request to {request_parameters.url!r} timed out after " - f"{request_parameters.timeout}s: {e}" - ), - category=UiPathErrorCategory.USER, + raise ToolException( + f"HTTP request timed out: Request to {request_parameters.url!r} " + f"timed out after {request_parameters.timeout}s: {e}" ) from e except httpx.HTTPError as e: - raise AgentRuntimeError( - code=AgentRuntimeErrorCode.HTTP_ERROR, - title="HTTP request failed", - detail=f"Request to {request_parameters.url!r} failed: {e}", - category=UiPathErrorCategory.USER, + raise ToolException( + f"HTTP request failed: Request to {request_parameters.url!r} " + f"failed: {e}" ) from e # Non-2xx responses are returned to the agent rather than raised, so it @@ -424,13 +445,14 @@ async def http_request_tool_fn(**kwargs: Any) -> dict[str, Any]: job_attachment_wrapper = get_job_attachment_wrapper(output_type=output_model) - tool = StructuredToolWithArgumentProperties( + tool = _HttpRequestTool( name=tool_name, description=resource.description, args_schema=input_model, coroutine=http_request_tool_fn, output_type=output_model, argument_properties=resource.argument_properties, + handle_tool_error=True, metadata={ "tool_type": resource.type.lower(), "display_name": tool_name, diff --git a/tests/agent/tools/internal_tools/test_http_request_tool.py b/tests/agent/tools/internal_tools/test_http_request_tool.py index 60d8b5031..ccc63dc83 100644 --- a/tests/agent/tools/internal_tools/test_http_request_tool.py +++ b/tests/agent/tools/internal_tools/test_http_request_tool.py @@ -6,16 +6,19 @@ without any DNS lookup — no need to mock internal functions. """ +from typing import Any from unittest.mock import AsyncMock, patch +import httpx import pytest +from langchain_core.messages import AIMessage, ToolMessage from pydantic import BaseModel from uipath.agent.models.agent import ( AgentInternalHttpRequestToolProperties, AgentInternalToolResourceConfig, ) -from uipath_langchain.agent.exceptions import AgentRuntimeError +from uipath_langchain.agent.react.types import AgentGraphState from uipath_langchain.agent.tools.internal_tools.http_request_tool import ( create_http_request_tool, ) @@ -91,6 +94,21 @@ def tool(resource_config, mock_llm): return create_http_request_tool(resource_config, mock_llm) +async def _call(tool: Any, args: dict[str, Any]) -> ToolMessage: + """Invoke the tool with a tool call, as the agent graph does.""" + result = await tool.ainvoke( + {"name": tool.name, "args": args, "id": "call_1", "type": "tool_call"} + ) + assert isinstance(result, ToolMessage) + return result + + +def _assert_returned_to_model(message: ToolMessage, expected: str) -> None: + """The error came back as an error tool message rather than being raised.""" + assert message.status == "error" + assert expected in message.content + + class TestCreateHttpRequestTool: def test_tool_creation_and_schema(self, tool): """Tool is created with the canonical input fields and fixed output fields.""" @@ -228,11 +246,41 @@ async def test_schemeless_url_defaults_to_https(self, tool, httpx_mock): assert result["statusCode"] == 200 assert str(httpx_mock.get_requests()[0].url) == BASE_URL - async def test_missing_url_raises(self, mock_llm): + async def test_pair_keys_are_matched_case_insensitively(self, tool, httpx_mock): + """Models sometimes write ``Name``/``Value``; the pair keys are lowercased. + + Only the keys are normalized: header and param values keep their case. + """ + httpx_mock.add_response(method="GET", status_code=200, text="ok") + + message = await _call( + tool, + { + "url": f"{BASE_URL}/x", + "headers": [{"Name": "Accept", "Value": "text/html"}], + "params": [{"NAME": "q", "value": "MixedCase"}], + }, + ) + + assert message.status == "success" + request = httpx_mock.get_requests()[0] + assert request.headers["accept"] == "text/html" + assert dict(request.url.params) == {"q": "MixedCase"} + + async def test_schema_mismatch_is_returned_to_model(self, tool): + """Arguments that still fail the input schema go back to the model.""" + message = await _call( + tool, {"url": f"{BASE_URL}/x", "headers": [{"key": "Accept"}]} + ) + + _assert_returned_to_model(message, "Invalid input for tool 'http_request'") + assert "headers.0.name" in message.content + + async def test_missing_url_is_returned_to_model(self, mock_llm): """The tool's own guard rejects a missing url. Uses a schema that does not mark ``url`` required, so args validation - passes and the tool's defensive check is what raises (a schema that + passes and the tool's defensive check is what reports it (a schema that requires ``url`` would be rejected earlier, by validation). """ resource = AgentInternalToolResourceConfig( @@ -243,11 +291,13 @@ async def test_missing_url_raises(self, mock_llm): properties=AgentInternalHttpRequestToolProperties(), ) tool = create_http_request_tool(resource, mock_llm) - with pytest.raises(AgentRuntimeError, match="Argument 'url' is required"): - await tool.ainvoke({}) - async def test_non_string_url_raises(self, mock_llm): - """A non-string url is rejected with a clean error. + message = await _call(tool, {}) + + _assert_returned_to_model(message, "Argument 'url' is required") + + async def test_non_string_url_is_returned_to_model(self, mock_llm): + """A non-string url is reported with a clean error. Uses a schema where ``url`` is untyped so validation passes a number through to the tool's own type guard. @@ -260,17 +310,21 @@ async def test_non_string_url_raises(self, mock_llm): properties=AgentInternalHttpRequestToolProperties(), ) tool = create_http_request_tool(resource, mock_llm) - with pytest.raises(AgentRuntimeError, match="'url' must be a string"): - await tool.ainvoke({"url": 123}) - async def test_invalid_method_raises(self, tool): - with pytest.raises(AgentRuntimeError, match="Unsupported HTTP method"): - await tool.ainvoke({"url": f"{BASE_URL}/x", "method": "FETCH"}) + message = await _call(tool, {"url": 123}) + + _assert_returned_to_model(message, "'url' must be a string") + + async def test_invalid_method_is_returned_to_model(self, tool): + message = await _call(tool, {"url": f"{BASE_URL}/x", "method": "FETCH"}) + + _assert_returned_to_model(message, "Unsupported HTTP method") @pytest.mark.parametrize("bad_timeout", [-1, 0, -0.5]) - async def test_non_positive_timeout_raises(self, tool, bad_timeout): - with pytest.raises(AgentRuntimeError, match="must be a positive number"): - await tool.ainvoke({"url": f"{BASE_URL}/x", "timeout": bad_timeout}) + async def test_non_positive_timeout_is_returned_to_model(self, tool, bad_timeout): + message = await _call(tool, {"url": f"{BASE_URL}/x", "timeout": bad_timeout}) + + _assert_returned_to_model(message, "must be a positive number") @pytest.mark.parametrize( "url", @@ -287,7 +341,60 @@ async def test_ssrf_blocked_targets_are_rejected(self, tool, url): """Internal/metadata/non-http targets are rejected before any request. No response is registered because the request never leaves the SSRF - guard. + guard. The refusal goes back to the model, which can pick another URL. + """ + message = await _call(tool, {"url": url}) + + assert message.status == "error" + assert ( + "Request was blocked" in message.content + or "Host could not be resolved" in message.content + ) + + @pytest.mark.parametrize( + ("exception", "expected"), + [ + (httpx.ReadTimeout("read timed out"), "HTTP request timed out"), + (httpx.ConnectError("connection refused"), "HTTP request failed"), + ], + ) + async def test_transport_errors_are_returned_to_model( + self, tool, httpx_mock, exception, expected + ): + httpx_mock.add_exception(exception) + + message = await _call(tool, {"url": f"{BASE_URL}/slow"}) + + _assert_returned_to_model(message, expected) + assert f"{BASE_URL}/slow" in message.content + + async def test_bad_call_does_not_fault_the_tool_node(self, tool): + """Through the graph's tool node, a bad call yields an error tool message. + + This is the non-conversational path, which has no node-level error + handling: before, the exception faulted the whole run. """ - with pytest.raises(AgentRuntimeError): - await tool.ainvoke({"url": url}) + from uipath_langchain.agent.tools import create_tool_node + + node = create_tool_node([tool])[tool.name] + state = AgentGraphState( + messages=[ + AIMessage( + content="", + tool_calls=[ + { + "name": tool.name, + "args": {"url": f"{BASE_URL}/x", "headers": [{"k": "v"}]}, + "id": "call_1", + "type": "tool_call", + } + ], + ) + ] + ) + + result = await node.ainvoke(state) + + message = result.update["messages"][0] + _assert_returned_to_model(message, "Invalid input for tool 'http_request'") + assert message.tool_call_id == "call_1"