Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_anthropic.py609 linesDownload Raw Back to wrappers
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 
codekingpro/portable-devtools · Team Ai