diff --git a/libs/sdk-py/langgraph_sdk/auth/__init__.py b/libs/sdk-py/langgraph_sdk/auth/__init__.py index f54f56706..54a0df989 100644 --- a/libs/sdk-py/langgraph_sdk/auth/__init__.py +++ b/libs/sdk-py/langgraph_sdk/auth/__init__.py @@ -728,7 +728,7 @@ def is_studio_user( return ( isinstance(user, types.StudioUser) or isinstance(user, dict) - and user.get("kind") == "StudioUser" + and user.get("kind") == "StudioUser" # ty: ignore[invalid-argument-type] ) diff --git a/libs/sdk-py/langgraph_sdk/client.py b/libs/sdk-py/langgraph_sdk/client.py index 2d1b5e082..6b8bafa55 100644 --- a/libs/sdk-py/langgraph_sdk/client.py +++ b/libs/sdk-py/langgraph_sdk/client.py @@ -437,6 +437,54 @@ class HttpClient: logger.error(f"Error from langgraph-api: {body}", exc_info=e) raise e + async def request_reconnect( + self, + path: str, + method: str, + *, + json: dict[str, Any] | None = None, + params: QueryParamTypes | None = None, + headers: Mapping[str, str] | None = None, + on_response: Callable[[httpx.Response], None] | None = None, + reconnect_limit: int = 5, + ) -> Any: + """Send a request that automatically reconnects to Location header.""" + request_headers, content = await _aencode_json(json) + if headers: + request_headers.update(headers) + async with self.client.stream( + method, path, headers=request_headers, content=content, params=params + ) as r: + if on_response: + on_response(r) + try: + r.raise_for_status() + except httpx.HTTPStatusError as e: + body = (await r.aread()).decode() + if sys.version_info >= (3, 11): + e.add_note(body) + else: + logger.error(f"Error from langgraph-api: {body}", exc_info=e) + raise e + loc = r.headers.get("location") + if reconnect_limit <= 0 or not loc: + return await _adecode_json(r) + try: + return await _adecode_json(r) + except httpx.HTTPError: + warnings.warn( + f"Request failed, attempting reconnect to Location: {loc}", + stacklevel=2, + ) + await r.aclose() + return await self.request_reconnect( + loc, + "GET", + headers=request_headers, + # don't pass on_response so it's only called once + reconnect_limit=reconnect_limit - 1, + ) + async def stream( self, path: str, @@ -2533,8 +2581,9 @@ class RunsClient: if on_run_created and (metadata := _get_run_metadata_from_response(res)): on_run_created(metadata) - response = await self.http.post( + response = await self.http.request_reconnect( endpoint, + "POST", json={k: v for k, v in payload.items() if v is not None}, params=params, headers=headers, @@ -2679,12 +2728,20 @@ class RunsClient: } if params: query_params.update(params) - return await self.http.post( - f"/threads/{thread_id}/runs/{run_id}/cancel", - json=None, - params=query_params, - headers=headers, - ) + if wait: + return await self.http.request_reconnect( + f"/threads/{thread_id}/runs/{run_id}/cancel", + "POST", + params=query_params, + headers=headers, + ) + else: + return await self.http.post( + f"/threads/{thread_id}/runs/{run_id}/cancel", + json=None, + params=query_params, + headers=headers, + ) async def join( self, @@ -2716,8 +2773,11 @@ class RunsClient: ``` """ # noqa: E501 - return await self.http.get( - f"/threads/{thread_id}/runs/{run_id}/join", headers=headers, params=params + return await self.http.request_reconnect( + f"/threads/{thread_id}/runs/{run_id}/join", + "GET", + headers=headers, + params=params, ) def join_stream( @@ -3689,6 +3749,54 @@ class SyncHttpClient: logger.error(f"Error from langgraph-api: {body}", exc_info=e) raise e + def request_reconnect( + self, + path: str, + method: str, + *, + json: dict[str, Any] | None = None, + params: QueryParamTypes | None = None, + headers: Mapping[str, str] | None = None, + on_response: Callable[[httpx.Response], None] | None = None, + reconnect_limit: int = 5, + ) -> Any: + """Send a request that automatically reconnects to Location header.""" + request_headers, content = _encode_json(json) + if headers: + request_headers.update(headers) + with self.client.stream( + method, path, headers=request_headers, content=content, params=params + ) as r: + if on_response: + on_response(r) + try: + r.raise_for_status() + except httpx.HTTPStatusError as e: + body = r.read().decode() + if sys.version_info >= (3, 11): + e.add_note(body) + else: + logger.error(f"Error from langgraph-api: {body}", exc_info=e) + raise e + loc = r.headers.get("location") + if reconnect_limit <= 0 or not loc: + return _decode_json(r) + try: + return _decode_json(r) + except httpx.HTTPError: + warnings.warn( + f"Request failed, attempting reconnect to Location: {loc}", + stacklevel=2, + ) + r.close() + return self.request_reconnect( + loc, + "GET", + headers=request_headers, + # don't pass on_response so it's only called once + reconnect_limit=reconnect_limit - 1, + ) + def stream( self, path: str, @@ -5754,8 +5862,9 @@ class SyncRunsClient: endpoint = ( f"/threads/{thread_id}/runs/wait" if thread_id is not None else "/runs/wait" ) - return self.http.post( + return self.http.request_reconnect( endpoint, + "POST", json={k: v for k, v in payload.items() if v is not None}, params=params, headers=headers, @@ -5878,11 +5987,25 @@ class SyncRunsClient: ``` """ # noqa: E501 + query_params = { + "wait": 1 if wait else 0, + "action": action, + } + if params: + query_params.update(params) + if wait: + return self.http.request_reconnect( + f"/threads/{thread_id}/runs/{run_id}/cancel", + "POST", + json=None, + params=query_params, + headers=headers, + ) return self.http.post( - f"/threads/{thread_id}/runs/{run_id}/cancel?wait={1 if wait else 0}&action={action}", + f"/threads/{thread_id}/runs/{run_id}/cancel", json=None, + params=query_params, headers=headers, - params=params, ) def join( @@ -5915,8 +6038,11 @@ class SyncRunsClient: ``` """ # noqa: E501 - return self.http.get( - f"/threads/{thread_id}/runs/{run_id}/join", headers=headers, params=params + return self.http.request_reconnect( + f"/threads/{thread_id}/runs/{run_id}/join", + "GET", + headers=headers, + params=params, ) def join_stream(