codekingpro/portable-devtools
114k
1"""Shared utility functions for async and sync clients."""2 3from __future__ import annotations4 5import functools6import os7import re8from collections.abc import Mapping9from datetime import tzinfo10from typing import TYPE_CHECKING, Any, cast11from urllib.parse import urlparse12 13import httpx14 15import langgraph_sdk16from langgraph_sdk.schema import RunCreateMetadata17 18if TYPE_CHECKING:19 from zoneinfo import ZoneInfo20 21RESERVED_HEADERS = ("x-api-key",)22 23NOT_PROVIDED = cast(None, object())24 25 26def _get_api_key(api_key: str | None = NOT_PROVIDED) -> str | None:27 """Get the API key from the environment.28 Precedence:29 1. explicit string argument30 2. LANGGRAPH_API_KEY (if api_key not provided)31 3. LANGSMITH_API_KEY (if api_key not provided)32 4. LANGCHAIN_API_KEY (if api_key not provided)33 34 Args:35 api_key: The API key to use. Can be:36 - A string: use this exact API key37 - None: explicitly skip loading from environment38 - NOT_PROVIDED (default): auto-load from environment variables39 """40 if isinstance(api_key, str):41 return api_key42 if api_key is NOT_PROVIDED:43 # api_key is not explicitly provided, try to load from environment44 for prefix in ["LANGGRAPH", "LANGSMITH", "LANGCHAIN"]:45 if env := os.getenv(f"{prefix}_API_KEY"):46 return env.strip().strip('"').strip("'")47 # api_key is explicitly None, don't load from environment48 return None49 50 51def _get_headers(52 api_key: str | None,53 custom_headers: Mapping[str, str] | None,54) -> dict[str, str]:55 """Combine api_key and custom user-provided headers."""56 custom_headers = custom_headers or {}57 for header in RESERVED_HEADERS:58 if header in custom_headers:59 raise ValueError(f"Cannot set reserved header '{header}'")60 61 headers = {62 "User-Agent": f"langgraph-sdk-py/{langgraph_sdk.__version__}",63 **custom_headers,64 }65 resolved_api_key = _get_api_key(api_key)66 if resolved_api_key:67 headers["x-api-key"] = resolved_api_key68 69 return headers70 71 72def _orjson_default(obj: Any) -> Any:73 is_class = isinstance(obj, type)74 if hasattr(obj, "model_dump") and callable(obj.model_dump):75 if is_class:76 raise TypeError(77 f"Cannot JSON-serialize type object: {obj!r}. Did you mean to pass an instance of the object instead?"78 f"\nReceived type: {obj!r}"79 )80 return obj.model_dump()81 elif hasattr(obj, "dict") and callable(obj.dict):82 if is_class:83 raise TypeError(84 f"Cannot JSON-serialize type object: {obj!r}. Did you mean to pass an instance of the object instead?"85 f"\nReceived type: {obj!r}"86 )87 return obj.dict()88 elif isinstance(obj, (set, frozenset)):89 return list(obj)90 else:91 raise TypeError(f"Object of type {type(obj)} is not JSON serializable")92 93 94# Compiled regex pattern for extracting run metadata from Content-Location header95_RUN_METADATA_PATTERN = re.compile(96 r"(\/threads\/(?P<thread_id>.+))?\/runs\/(?P<run_id>.+)"97)98 99 100def _get_run_metadata_from_response(101 response: httpx.Response,102) -> RunCreateMetadata | None:103 """Extract run metadata from the response headers."""104 if (content_location := response.headers.get("Content-Location")) and (105 match := _RUN_METADATA_PATTERN.search(content_location)106 ):107 return RunCreateMetadata(108 run_id=match.group("run_id"),109 thread_id=match.group("thread_id") or None,110 )111 112 return None113 114 115def _sse_to_v2_dict(event: str, data: Any) -> dict[str, Any] | None:116 """Convert an SSE event+data pair into a v2 stream part dict.117 118 Returns None for ``end`` events (signals end of stream).119 """120 if event == "end":121 return None122 parts = event.split("|")123 event_type = parts[0]124 ns = parts[1:] if len(parts) > 1 else []125 result: dict[str, Any] = {"type": event_type, "ns": ns, "data": data}126 if event_type == "values" and isinstance(data, dict):127 result["interrupts"] = data.pop("__interrupt__", [])128 else:129 result["interrupts"] = []130 return result131 132 133def _resolve_timezone(tz: str | tzinfo | ZoneInfo | None) -> str | None:134 """Convert a timezone argument to an IANA timezone string.135 136 Accepts:137 - A string (returned as-is, assumed to be an IANA timezone name)138 - A ``datetime.tzinfo`` instance (e.g. ``zoneinfo.ZoneInfo("America/New_York")``,139 ``datetime.timezone.utc``). The ``key`` attribute is used if available,140 otherwise ``tzname(None)`` is used.141 - ``None`` (returned as ``None``)142 """143 if tz is None or isinstance(tz, str):144 return tz145 if isinstance(tz, tzinfo):146 # ZoneInfo objects have a .key attribute with the IANA name147 if hasattr(tz, "key"):148 return tz.key # type: ignore[union-attr]149 # Fall back to tzname for fixed-offset timezones like datetime.timezone.utc150 name = tz.tzname(None)151 if name is not None:152 return name153 raise ValueError(154 f"Cannot determine timezone name from {tz!r}. "155 "Use a zoneinfo.ZoneInfo instance or pass a string like 'America/New_York'."156 )157 raise TypeError(158 f"Expected str, datetime.tzinfo, or None for timezone, got {type(tz).__name__}"159 )160 161 162def _default_port(scheme: str) -> int:163 return 443 if scheme == "https" else 80164 165 166def _validate_reconnect_location(base_url: httpx.URL, location: str) -> str:167 """Validate that a reconnect Location URL is same-origin as the base URL.168 169 Raises ValueError if the Location header points to a different origin170 (scheme + host + port), which would leak credentials to an external server.171 """172 parsed = urlparse(location)173 # Relative URLs are safe — they resolve against the base174 if not parsed.scheme and not parsed.netloc:175 return location176 # Compare origin components (normalize default ports to avoid mismatches)177 base_scheme = str(base_url.scheme)178 base_origin = (179 base_scheme,180 str(base_url.host),181 base_url.port or _default_port(base_scheme),182 )183 loc_origin = (184 parsed.scheme,185 parsed.hostname or "",186 parsed.port or _default_port(parsed.scheme),187 )188 if base_origin != loc_origin:189 raise ValueError(190 f"Refusing to follow cross-origin reconnect Location: {location!r} "191 f"(origin {loc_origin}) does not match base URL origin {base_origin}"192 )193 return location194 195 196def _provided_vals(d: Mapping[str, Any]) -> dict[str, Any]:197 return {k: v for k, v in d.items() if v is not None}198 199 200_registered_transports: list[httpx.ASGITransport] = []201 202 203# Do not move; this is used in the server.204def configure_loopback_transports(app: Any) -> None:205 for transport in _registered_transports:206 transport.app = app207 208 209@functools.lru_cache(maxsize=1)210def get_asgi_transport() -> type[httpx.ASGITransport]:211 try:212 from langgraph_api import asgi_transport # type: ignore[unresolved-import]213 214 return asgi_transport.ASGITransport215 except ImportError:216 # Older versions of the server217 return httpx.ASGITransport218 