codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import functools4import logging5import warnings6from collections.abc import AsyncIterator, Mapping, Sequence7from typing import (8 TYPE_CHECKING,9 Any,10 Callable,11 Optional,12 TypeVar,13 Union,14)15 16from typing_extensions import Self, TypedDict17 18from langsmith import client as ls_client19from langsmith import run_helpers20from langsmith._internal._orjson import dumps as _dumps21from langsmith.schemas import InputTokenDetails, UsageMetadata22 23if TYPE_CHECKING:24 import httpx25 from anthropic import Anthropic, AsyncAnthropic26 from anthropic.lib.streaming import AsyncMessageStream, MessageStream27 from anthropic.types import Completion, Message, MessageStreamEvent28 29C = TypeVar("C", bound=Union["Anthropic", "AsyncAnthropic", Any])30logger = logging.getLogger(__name__)31 32 33@functools.lru_cache34def _get_not_given() -> Optional[tuple[type, ...]]:35 try:36 from anthropic._types import NotGiven, Omit37 38 return (NotGiven, Omit)39 except ImportError:40 return None41 42 43def _strip_not_given(d: dict) -> dict:44 try:45 if not_given := _get_not_given():46 d = {47 k: v48 for k, v in d.items()49 if not any(isinstance(v, t) for t in not_given)50 }51 except Exception as e:52 logger.error(f"Error stripping NotGiven: {e}")53 54 if "system" in d:55 d["messages"] = [{"role": "system", "content": d["system"]}] + d.get(56 "messages", []57 )58 d.pop("system")59 return {k: v for k, v in d.items() if v is not None}60 61 62def _infer_ls_params(prepopulated_invocation_params: dict, kwargs: dict):63 stripped = _strip_not_given(kwargs)64 65 stop = stripped.get("stop")66 if stop and isinstance(stop, str):67 stop = [stop]68 69 # Allowlist of safe invocation parameters to include70 # Only include known, non-sensitive parameters71 allowed_invocation_keys = {72 "mcp_servers",73 "service_tier",74 "tool_choice",75 "top_k",76 "top_p",77 "stream",78 "thinking",79 }80 81 # Only include allowlisted parameters82 invocation_params = {83 k: v for k, v in stripped.items() if k in allowed_invocation_keys84 }85 86 return {87 "ls_provider": "anthropic",88 "ls_model_type": "chat",89 "ls_model_name": stripped.get("model", None),90 "ls_temperature": stripped.get("temperature", None),91 "ls_max_tokens": stripped.get("max_tokens", None),92 "ls_stop": stop,93 "ls_invocation_params": {94 **prepopulated_invocation_params,95 **invocation_params,96 },97 }98 99 100@functools.lru_cache101def _get_sdk_accumulate_event() -> Optional[Callable]:102 try:103 from anthropic.lib.streaming._messages import accumulate_event104 105 return accumulate_event106 except ImportError:107 return None108 109 110def _create_usage_metadata(anthropic_token_usage: dict) -> UsageMetadata:111 input_tokens = anthropic_token_usage.get("input_tokens") or 0112 output_tokens = anthropic_token_usage.get("output_tokens") or 0113 114 input_token_details: dict = {}115 cache_read = anthropic_token_usage.get("cache_read_input_tokens") or 0116 if cache_read:117 input_token_details["cache_read"] = cache_read118 119 cache_creation_obj = anthropic_token_usage.get("cache_creation") or {}120 if cache_creation_obj:121 ephemeral_5m = cache_creation_obj.get("ephemeral_5m_input_tokens") or 0122 ephemeral_1h = cache_creation_obj.get("ephemeral_1h_input_tokens") or 0123 if ephemeral_5m:124 input_token_details["ephemeral_5m_input_tokens"] = ephemeral_5m125 if ephemeral_1h:126 input_token_details["ephemeral_1h_input_tokens"] = ephemeral_1h127 else:128 cache_creation = anthropic_token_usage.get("cache_creation_input_tokens") or 0129 if cache_creation:130 input_token_details["cache_creation"] = cache_creation131 132 # Anthropic cache tokens are ADDITIVE (not subsets of input_tokens like OpenAI).133 # Sum them into input_tokens so the backend cost calculation is correct.134 cache_token_sum = sum(input_token_details.values())135 adjusted_input = input_tokens + cache_token_sum136 adjusted_total = adjusted_input + output_tokens137 138 result = UsageMetadata(139 input_tokens=adjusted_input,140 output_tokens=output_tokens,141 total_tokens=adjusted_total,142 )143 if input_token_details:144 result["input_token_details"] = InputTokenDetails(**input_token_details)145 return result146 147 148def _message_to_outputs(message: Any) -> dict:149 """Convert an Anthropic Message to a flat outputs dict with usage_metadata."""150 # ParsedBetaMessage/ParsedMessage (from beta.messages.parse()) carry user-defined151 # Pydantic models in parsed_output and ParsedBetaTextBlock in content. These trigger152 # PydanticSerializationUnexpectedValue warnings because the values do not match the153 # declared union types in the base BetaMessage schema. Suppress for parsed types.154 if hasattr(message, "parsed_output"):155 with warnings.catch_warnings():156 warnings.simplefilter("ignore")157 outputs = message.model_dump()158 else:159 outputs = message.model_dump()160 anthropic_token_usage = outputs.pop("usage", None)161 if anthropic_token_usage:162 outputs["usage_metadata"] = _create_usage_metadata(anthropic_token_usage)163 outputs.pop("type", None)164 165 content = outputs.get("content") or []166 tool_use_blocks = [167 b for b in content if isinstance(b, dict) and b.get("type") == "tool_use"168 ]169 if tool_use_blocks:170 text_parts = [171 b.get("text", "")172 for b in content173 if isinstance(b, dict) and b.get("type") == "text"174 ]175 outputs["content"] = "".join(text_parts) or None176 outputs["tool_calls"] = [177 {178 "id": block.get("id", f"call_{i}"),179 "type": "function",180 "index": i,181 "function": {182 "name": block.get("name", ""),183 "arguments": _dumps(block.get("input", {})).decode(),184 },185 }186 for i, block in enumerate(tool_use_blocks)187 ]188 return outputs189 190 191def _reduce_chat_chunks(all_chunks: Sequence) -> dict:192 accumulate = _get_sdk_accumulate_event()193 if accumulate is None:194 return {"output": all_chunks}195 full_message = None196 for chunk in all_chunks:197 try:198 full_message = accumulate(199 event=chunk,200 current_snapshot=full_message,201 )202 except RuntimeError as e:203 logger.debug(f"Error accumulating event in Anthropic Wrapper: {e}")204 return {"output": all_chunks}205 if full_message is None:206 return {"output": all_chunks}207 return _message_to_outputs(full_message)208 209 210def _reduce_completions(all_chunks: list[Completion]) -> dict:211 all_content = []212 for chunk in all_chunks:213 content = chunk.completion214 if content is not None:215 all_content.append(content)216 content = "".join(all_content)217 if all_chunks:218 d = all_chunks[-1].model_dump()219 d["choices"] = [{"text": content}]220 else:221 d = {"choices": [{"text": content}]}222 223 return d224 225 226def _process_chat_completion(outputs: Any):227 try:228 # Check if outputs is a LegacyAPIResponse wrapper (from with_raw_response).229 # The Anthropic SDK's LegacyAPIResponse wraps the actual response object.230 # Call .parse() to extract the Message for tracing.231 # See: anthropics/anthropic-sdk-python _legacy_response.py#L102232 if hasattr(outputs, "parse") and callable(outputs.parse):233 try:234 outputs = outputs.parse()235 except Exception:236 pass237 return _message_to_outputs(outputs)238 except BaseException as e:239 logger.debug(f"Error processing chat completion: {e}")240 return {"output": outputs}241 242 243def _get_wrapper(244 original_create: Callable,245 name: str,246 reduce_fn: Callable,247 prepopulated_invocation_params: dict,248 tracing_extra: TracingExtra,249) -> Callable:250 @functools.wraps(original_create)251 def create(*args, **kwargs):252 stream = kwargs.get("stream")253 decorator = run_helpers.traceable(254 name=name,255 run_type="llm",256 reduce_fn=reduce_fn if stream else None,257 process_inputs=_strip_not_given,258 process_outputs=_process_chat_completion,259 _invocation_params_fn=functools.partial(260 _infer_ls_params, prepopulated_invocation_params261 ),262 **tracing_extra,263 )264 265 result = decorator(original_create)(*args, **kwargs)266 return result267 268 @functools.wraps(original_create)269 async def acreate(*args, **kwargs):270 stream = kwargs.get("stream")271 decorator = run_helpers.traceable(272 name=name,273 run_type="llm",274 reduce_fn=reduce_fn if stream else None,275 process_inputs=_strip_not_given,276 process_outputs=_process_chat_completion,277 _invocation_params_fn=functools.partial(278 _infer_ls_params, prepopulated_invocation_params279 ),280 **tracing_extra,281 )282 result = await decorator(original_create)(*args, **kwargs)283 return result284 285 return acreate if run_helpers.is_async(original_create) else create286 287 288def _get_stream_wrapper(289 original_stream: Callable,290 name: str,291 prepopulated_invocation_params: dict,292 tracing_extra: TracingExtra,293) -> Callable:294 """Create a wrapper for Anthropic's streaming context manager."""295 is_async = "async" in str(original_stream).lower()296 configured_traceable = run_helpers.traceable(297 name=name,298 reduce_fn=_reduce_chat_chunks,299 run_type="llm",300 process_inputs=_strip_not_given,301 _invocation_params_fn=functools.partial(302 _infer_ls_params, prepopulated_invocation_params303 ),304 **tracing_extra,305 )306 configured_traceable_text = run_helpers.traceable(307 name=name,308 run_type="llm",309 process_inputs=_strip_not_given,310 process_outputs=_process_chat_completion,311 _invocation_params_fn=functools.partial(312 _infer_ls_params, prepopulated_invocation_params313 ),314 **tracing_extra,315 )316 317 if is_async:318 319 class AsyncMessageStreamWrapper:320 def __init__(321 self,322 wrapped: AsyncMessageStream,323 **kwargs,324 ) -> None:325 self._wrapped = wrapped326 self._kwargs = kwargs327 328 @property329 def text_stream(self):330 @configured_traceable_text331 async def _text_stream(**_):332 async for chunk in self._wrapped.text_stream:333 yield chunk334 run_tree = run_helpers.get_current_run_tree()335 final_message = await self._wrapped.get_final_message()336 outputs = _message_to_outputs(final_message)337 run_tree.outputs = outputs338 if usage := outputs.get("usage_metadata"):339 run_tree.metadata["usage_metadata"] = usage340 341 return _text_stream(**self._kwargs)342 343 @property344 def response(self) -> httpx.Response:345 return self._wrapped.response346 347 @property348 def request_id(self) -> str | None:349 return self._wrapped.request_id350 351 async def __anext__(self) -> MessageStreamEvent:352 aiter = self.__aiter__()353 return await aiter.__anext__()354 355 async def __aiter__(self) -> AsyncIterator[MessageStreamEvent]:356 @configured_traceable357 def traced_iter(**_):358 return self._wrapped.__aiter__()359 360 async for chunk in traced_iter(**self._kwargs):361 yield chunk362 363 async def __aenter__(self) -> Self:364 await self._wrapped.__aenter__()365 return self366 367 async def __aexit__(self, *exc) -> None:368 await self._wrapped.__aexit__(*exc)369 370 async def close(self) -> None:371 await self._wrapped.close()372 373 async def get_final_message(self) -> Message:374 return await self._wrapped.get_final_message()375 376 async def get_final_text(self) -> str:377 return await self._wrapped.get_final_text()378 379 async def until_done(self) -> None:380 await self._wrapped.until_done()381 382 @property383 def current_message_snapshot(self) -> Message:384 return self._wrapped.current_message_snapshot385 386 class AsyncMessagesStreamManagerWrapper:387 def __init__(self, **kwargs):388 self._kwargs = kwargs389 390 async def __aenter__(self):391 self._manager = original_stream(**self._kwargs)392 stream = await self._manager.__aenter__()393 return AsyncMessageStreamWrapper(stream, **self._kwargs)394 395 async def __aexit__(self, *exc):396 await self._manager.__aexit__(*exc)397 398 return AsyncMessagesStreamManagerWrapper399 else:400 401 class MessageStreamWrapper:402 def __init__(403 self,404 wrapped: MessageStream,405 **kwargs,406 ) -> None:407 self._wrapped = wrapped408 self._kwargs = kwargs409 410 @property411 def response(self) -> Any:412 return self._wrapped.response413 414 @property415 def request_id(self) -> str | None:416 return self._wrapped.request_id # type: ignore[no-any-return]417 418 @property419 def text_stream(self):420 @configured_traceable_text421 def _text_stream(**_):422 yield from self._wrapped.text_stream423 run_tree = run_helpers.get_current_run_tree()424 final_message = self._wrapped.get_final_message()425 outputs = _message_to_outputs(final_message)426 run_tree.outputs = outputs427 if usage := outputs.get("usage_metadata"):428 run_tree.metadata["usage_metadata"] = usage429 430 return _text_stream(**self._kwargs)431 432 def __next__(self) -> MessageStreamEvent:433 return self.__iter__().__next__()434 435 def __iter__(self):436 @configured_traceable437 def traced_iter(**_):438 return self._wrapped.__iter__()439 440 return traced_iter(**self._kwargs)441 442 def __enter__(self) -> Self:443 self._wrapped.__enter__()444 return self445 446 def __exit__(self, *exc) -> None:447 self._wrapped.__exit__(*exc)448 449 def close(self) -> None:450 self._wrapped.close()451 452 def get_final_message(self) -> Message:453 return self._wrapped.get_final_message()454 455 def get_final_text(self) -> str:456 return self._wrapped.get_final_text()457 458 def until_done(self) -> None:459 return self._wrapped.until_done()460 461 @property462 def current_message_snapshot(self) -> Message:463 return self._wrapped.current_message_snapshot464 465 class MessagesStreamManagerWrapper:466 def __init__(self, **kwargs):467 self._kwargs = kwargs468 469 def __enter__(self):470 self._manager = original_stream(**self._kwargs)471 return MessageStreamWrapper(self._manager.__enter__(), **self._kwargs)472 473 def __exit__(self, *exc):474 self._manager.__exit__(*exc)475 476 return MessagesStreamManagerWrapper477 478 479class TracingExtra(TypedDict, total=False):480 metadata: Optional[Mapping[str, Any]]481 tags: Optional[list[str]]482 client: Optional[ls_client.Client]483 484 485def wrap_anthropic(486 client: C,487 *,488 tracing_extra: Optional[TracingExtra] = None,489 chat_name: str = "ChatAnthropic",490 completions_name: str = "Anthropic",491) -> C:492 """Patch the Anthropic client to make it traceable.493 494 Args:495 client: The client to patch.496 tracing_extra: Extra tracing information.497 chat_name: The run name for the messages endpoint.498 completions_name: The run name for the completions endpoint.499 500 Returns:501 The patched client.502 503 Example:504 ```python505 import anthropic506 from langsmith import wrappers507 508 client = wrappers.wrap_anthropic(anthropic.Anthropic())509 510 # Use Anthropic client same as you normally would:511 system = "You are a helpful assistant."512 messages = [513 {514 "role": "user",515 "content": "What physics breakthroughs do you predict will happen by 2300?",516 }517 ]518 completion = client.messages.create(519 model="claude-3-5-sonnet-latest",520 messages=messages,521 max_tokens=1000,522 system=system,523 )524 print(completion.content)525 526 # With raw response to access headers:527 raw_response = client.messages.with_raw_response.create(528 model="claude-3-5-sonnet-latest",529 messages=messages,530 max_tokens=1000,531 system=system,532 )533 print(raw_response.headers) # Access HTTP headers534 message = raw_response.parse() # Get parsed response535 536 # You can also use the streaming context manager:537 with client.messages.stream(538 model="claude-3-5-sonnet-latest",539 messages=messages,540 max_tokens=1000,541 system=system,542 ) as stream:543 for text in stream.text_stream:544 print(text, end="", flush=True)545 message = stream.get_final_message()546 ```547 """ # noqa: E501548 tracing_extra = tracing_extra or {}549 550 # Extract ls_invocation_params from metadata551 metadata = dict(tracing_extra.get("metadata") or {})552 prepopulated_invocation_params = metadata.pop("ls_invocation_params", {})553 554 # Create new tracing_extra without ls_invocation_params in metadata555 tracing_extra_rest: TracingExtra = { # type: ignore[assignment]556 k: v for k, v in tracing_extra.items() if k != "metadata"557 }558 if metadata:559 tracing_extra_rest["metadata"] = metadata # type: ignore[typeddict-item]560 561 client.messages.create = _get_wrapper( # type: ignore[method-assign]562 client.messages.create,563 chat_name,564 _reduce_chat_chunks,565 prepopulated_invocation_params,566 tracing_extra_rest,567 )568 569 client.messages.stream = _get_stream_wrapper( # type: ignore[method-assign]570 client.messages.stream,571 chat_name,572 prepopulated_invocation_params,573 tracing_extra_rest,574 )575 client.completions.create = _get_wrapper( # type: ignore[method-assign]576 client.completions.create,577 completions_name,578 _reduce_completions,579 prepopulated_invocation_params,580 tracing_extra_rest,581 )582 583 if (584 hasattr(client, "beta")585 and hasattr(client.beta, "messages")586 and hasattr(client.beta.messages, "create")587 ):588 client.beta.messages.create = _get_wrapper( # type: ignore[method-assign]589 client.beta.messages.create, # type: ignore590 chat_name,591 _reduce_chat_chunks,592 prepopulated_invocation_params,593 tracing_extra_rest,594 )595 596 if (597 hasattr(client, "beta")598 and hasattr(client.beta, "messages")599 and hasattr(client.beta.messages, "parse")600 ):601 client.beta.messages.parse = _get_wrapper( # type: ignore[method-assign]602 client.beta.messages.parse, # type: ignore603 chat_name,604 _reduce_chat_chunks,605 prepopulated_invocation_params,606 tracing_extra_rest,607 )608 return client609 