Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_openai.py649 linesDownload Raw Back to wrappers
1from __future__ import annotations2 3import functools4import logging5from collections import defaultdict6from collections.abc import Mapping7from typing import (8    TYPE_CHECKING,9    Any,10    Callable,11    Optional,12    TypeVar,13    Union,14)15 16from typing_extensions import TypedDict17 18from langsmith import client as ls_client19from langsmith import run_helpers20from langsmith.schemas import InputTokenDetails, OutputTokenDetails, UsageMetadata21 22if TYPE_CHECKING:23    from openai import AsyncOpenAI, OpenAI24    from openai.types.chat.chat_completion_chunk import (25        ChatCompletionChunk,26        Choice,27        ChoiceDeltaToolCall,28    )29    from openai.types.completion import Completion30    from openai.types.responses import ResponseStreamEvent  # type: ignore31 32# Any is used since it may work with Azure or other providers33C = TypeVar("C", bound=Union["OpenAI", "AsyncOpenAI", Any])34logger = logging.getLogger(__name__)35 36 37@functools.lru_cache38def _get_omit_types() -> tuple[type, ...]:39    """Get NotGiven/Omit sentinel types used by OpenAI SDK."""40    types: list[type[Any]] = []41    try:42        from openai._types import NotGiven, Omit43 44        types.append(NotGiven)45        types.append(Omit)46    except ImportError:47        pass48 49    return tuple(types)50 51 52def _strip_not_given(d: dict) -> dict:53    try:54        omit_types = _get_omit_types()55        if not omit_types:56            return d57        return {58            k: v59            for k, v in d.items()60            if not (isinstance(v, omit_types) or (k.startswith("extra_") and v is None))61        }62    except Exception as e:63        logger.error(f"Error stripping NotGiven: {e}")64        return d65 66 67def _process_inputs(d: dict) -> dict:68    """Strip `NotGiven` values and serialize `text_format` to JSON schema."""69    d = _strip_not_given(d)70 71    # Convert text_format (Pydantic model) to JSON schema if present72    if "text_format" in d:73        text_format = d["text_format"]74        if hasattr(text_format, "model_json_schema"):75            try:76                return {77                    **d,78                    "text_format": text_format.model_json_schema(),79                }80            except Exception:81                pass82    return d83 84 85def _infer_invocation_params(86    model_type: str,87    provider: str,88    prepopulated_invocation_params: dict,89    use_responses_api: bool,90    kwargs: dict,91):92    stripped = _strip_not_given(kwargs)93 94    stop = stripped.get("stop")95    if stop and isinstance(stop, str):96        stop = [stop]97 98    # Allowlist of safe invocation parameters to include99    # Only include known, non-sensitive parameters100    allowed_invocation_keys = {101        "frequency_penalty",102        "n",103        "logit_bias",104        "logprobs",105        "modalities",106        "parallel_tool_calls",107        "prediction",108        "presence_penalty",109        "prompt_cache_key",110        "reasoning",111        "reasoning_effort",112        "response_format",113        "seed",114        "service_tier",115        "stream_options",116        "top_logprobs",117        "top_p",118        "truncation",119        "user",120        "verbosity",121        "web_search_options",122    }123 124    # Only include allowlisted parameters125    invocation_params = {126        k: v for k, v in stripped.items() if k in allowed_invocation_keys127    }128 129    if use_responses_api:130        invocation_params["use_responses_api"] = True131 132    return {133        "ls_provider": provider,134        "ls_model_type": model_type,135        "ls_model_name": stripped.get("model"),136        "ls_temperature": stripped.get("temperature"),137        "ls_max_tokens": stripped.get("max_tokens")138        or stripped.get("max_completion_tokens")139        or stripped.get("max_output_tokens"),140        "ls_stop": stop,141        "ls_invocation_params": {142            **prepopulated_invocation_params,143            **invocation_params,144        },145    }146 147 148def _reduce_choices(choices: list[Choice]) -> dict:149    reversed_choices = list(reversed(choices))150    message: dict[str, Any] = {151        "role": "assistant",152        "content": "",153    }154    for c in reversed_choices:155        if hasattr(c, "delta") and getattr(c.delta, "role", None):156            message["role"] = c.delta.role157            break158    tool_calls: defaultdict[int, list[ChoiceDeltaToolCall]] = defaultdict(list)159    for c in choices:160        if hasattr(c, "delta"):161            if getattr(c.delta, "content", None):162                message["content"] += c.delta.content163            if getattr(c.delta, "function_call", None):164                if not message.get("function_call"):165                    message["function_call"] = {"name": "", "arguments": ""}166                name_ = getattr(c.delta.function_call, "name", None)167                if name_:168                    message["function_call"]["name"] += name_169                arguments_ = getattr(c.delta.function_call, "arguments", None)170                if arguments_:171                    message["function_call"]["arguments"] += arguments_172            if getattr(c.delta, "tool_calls", None):173                tool_calls_list = c.delta.tool_calls174                if tool_calls_list is not None:175                    for tool_call in tool_calls_list:176                        tool_calls[tool_call.index].append(tool_call)177    if tool_calls:178        message["tool_calls"] = [None for _ in range(max(tool_calls.keys()) + 1)]179        for index, tool_call_chunks in tool_calls.items():180            message["tool_calls"][index] = {181                "index": index,182                "id": next((c.id for c in tool_call_chunks if c.id), None),183                "type": next((c.type for c in tool_call_chunks if c.type), None),184                "function": {"name": "", "arguments": ""},185            }186            for chunk in tool_call_chunks:187                if getattr(chunk, "function", None):188                    name_ = getattr(chunk.function, "name", None)189                    if name_:190                        message["tool_calls"][index]["function"]["name"] += name_191                    arguments_ = getattr(chunk.function, "arguments", None)192                    if arguments_:193                        message["tool_calls"][index]["function"]["arguments"] += (194                            arguments_195                        )196    return {197        "index": getattr(choices[0], "index", 0) if choices else 0,198        "finish_reason": next(199            (200                c.finish_reason201                for c in reversed_choices202                if getattr(c, "finish_reason", None)203            ),204            None,205        ),206        "message": message,207    }208 209 210def _reduce_chat(all_chunks: list[ChatCompletionChunk]) -> dict:211    choices_by_index: defaultdict[int, list[Choice]] = defaultdict(list)212    for chunk in all_chunks:213        for choice in chunk.choices:214            choices_by_index[choice.index].append(choice)215    if all_chunks:216        d = all_chunks[-1].model_dump()217        d["choices"] = [218            _reduce_choices(choices) for choices in choices_by_index.values()219        ]220    else:221        d = {"choices": [{"message": {"role": "assistant", "content": ""}}]}222    # streamed outputs don't go through `process_outputs`223    # so we need to flatten metadata here224    oai_token_usage = d.pop("usage", None)225    d["usage_metadata"] = (226        _create_usage_metadata(oai_token_usage) if oai_token_usage else None227    )228    return d229 230 231def _reduce_completions(all_chunks: list[Completion]) -> dict:232    all_content = []233    for chunk in all_chunks:234        content = chunk.choices[0].text235        if content is not None:236            all_content.append(content)237    content = "".join(all_content)238    if all_chunks:239        d = all_chunks[-1].model_dump()240        d["choices"] = [{"text": content}]241    else:242        d = {"choices": [{"text": content}]}243 244    return d245 246 247def _create_usage_metadata(248    oai_token_usage: dict, service_tier: Optional[str] = None249) -> UsageMetadata:250    recognized_service_tier = (251        service_tier if service_tier in ["priority", "flex"] else None252    )253    service_tier_prefix = (254        f"{recognized_service_tier}_" if recognized_service_tier else ""255    )256 257    input_tokens = (258        oai_token_usage.get("prompt_tokens") or oai_token_usage.get("input_tokens") or 0259    )260    output_tokens = (261        oai_token_usage.get("completion_tokens")262        or oai_token_usage.get("output_tokens")263        or 0264    )265    total_tokens = oai_token_usage.get("total_tokens") or input_tokens + output_tokens266    input_token_details: dict = {267        "audio": (268            oai_token_usage.get("prompt_tokens_details")269            or oai_token_usage.get("input_tokens_details")270            or {}271        ).get("audio_tokens"),272        f"{service_tier_prefix}cache_read": (273            oai_token_usage.get("prompt_tokens_details")274            or oai_token_usage.get("input_tokens_details")275            or {}276        ).get("cached_tokens"),277    }278    output_token_details: dict = {279        "audio": (280            oai_token_usage.get("completion_tokens_details")281            or oai_token_usage.get("output_tokens_details")282            or {}283        ).get("audio_tokens"),284        f"{service_tier_prefix}reasoning": (285            oai_token_usage.get("completion_tokens_details")286            or oai_token_usage.get("output_tokens_details")287            or {}288        ).get("reasoning_tokens"),289    }290 291    if recognized_service_tier:292        # Avoid counting cache read and reasoning tokens towards the293        # service tier token count since service tier tokens are already294        # priced differently295        input_token_details[recognized_service_tier] = input_tokens - (296            input_token_details.get(f"{service_tier_prefix}cache_read") or 0297        )298        output_token_details[recognized_service_tier] = output_tokens - (299            output_token_details.get(f"{service_tier_prefix}reasoning") or 0300        )301 302    return UsageMetadata(303        input_tokens=input_tokens,304        output_tokens=output_tokens,305        total_tokens=total_tokens,306        input_token_details=InputTokenDetails(307            **{k: v for k, v in input_token_details.items() if v is not None}308        ),309        output_token_details=OutputTokenDetails(310            **{k: v for k, v in output_token_details.items() if v is not None}311        ),312    )313 314 315def _process_chat_completion(outputs: Any):316    try:317        # Check if outputs is an APIResponse wrapper (from with_raw_response).318        # The OpenAI SDK's APIResponse wraps the actual response object.319        # Call .parse() to extract the ChatCompletion/Completion for tracing.320        # See: github.com/openai/openai-python/blob/main/src/openai/_response.py#L285321        if hasattr(outputs, "parse") and callable(outputs.parse):322            try:323                outputs = outputs.parse()324            except Exception:325                pass326 327        rdict = outputs.model_dump()328        oai_token_usage = rdict.pop("usage", None)329        rdict["usage_metadata"] = (330            _create_usage_metadata(oai_token_usage, rdict.get("service_tier"))331            if oai_token_usage332            else None333        )334        return rdict335    except BaseException as e:336        logger.debug(f"Error processing chat completion: {e}")337        return {"output": outputs}338 339 340def _get_wrapper(341    original_create: Callable,342    name: str,343    reduce_fn: Callable,344    tracing_extra: Optional[TracingExtra] = None,345    invocation_params_fn: Optional[Callable] = None,346    process_outputs: Optional[Callable] = None,347) -> Callable:348    textra = tracing_extra or {}349 350    @functools.wraps(original_create)351    def create(*args, **kwargs):352        decorator = run_helpers.traceable(353            name=name,354            run_type="llm",355            reduce_fn=reduce_fn if kwargs.get("stream") is True else None,356            process_inputs=_process_inputs,357            _invocation_params_fn=invocation_params_fn,358            process_outputs=process_outputs,359            **textra,360        )361 362        return decorator(original_create)(*args, **kwargs)363 364    @functools.wraps(original_create)365    async def acreate(*args, **kwargs):366        decorator = run_helpers.traceable(367            name=name,368            run_type="llm",369            reduce_fn=reduce_fn if kwargs.get("stream") is True else None,370            process_inputs=_process_inputs,371            _invocation_params_fn=invocation_params_fn,372            process_outputs=process_outputs,373            **textra,374        )375        return await decorator(original_create)(*args, **kwargs)376 377    return acreate if run_helpers.is_async(original_create) else create378 379 380def _get_parse_wrapper(381    original_parse: Callable,382    name: str,383    process_outputs: Callable,384    tracing_extra: Optional[TracingExtra] = None,385    invocation_params_fn: Optional[Callable] = None,386) -> Callable:387    textra = tracing_extra or {}388 389    @functools.wraps(original_parse)390    def parse(*args, **kwargs):391        decorator = run_helpers.traceable(392            name=name,393            run_type="llm",394            reduce_fn=None,395            process_inputs=_process_inputs,396            _invocation_params_fn=invocation_params_fn,397            process_outputs=process_outputs,398            **textra,399        )400        return decorator(original_parse)(*args, **kwargs)401 402    @functools.wraps(original_parse)403    async def aparse(*args, **kwargs):404        decorator = run_helpers.traceable(405            name=name,406            run_type="llm",407            reduce_fn=None,408            process_inputs=_process_inputs,409            _invocation_params_fn=invocation_params_fn,410            process_outputs=process_outputs,411            **textra,412        )413        return await decorator(original_parse)(*args, **kwargs)414 415    return aparse if run_helpers.is_async(original_parse) else parse416 417 418def _reduce_response_events(events: list[ResponseStreamEvent]) -> dict:419    for event in events:420        if event.type == "response.completed":421            return _process_responses_api_output(event.response)422    return {}423 424 425class TracingExtra(TypedDict, total=False):426    metadata: Optional[Mapping[str, Any]]427    tags: Optional[list[str]]428    client: Optional[ls_client.Client]429 430 431def wrap_openai(432    client: C,433    *,434    tracing_extra: Optional[TracingExtra] = None,435    chat_name: str = "ChatOpenAI",436    completions_name: str = "OpenAI",437) -> C:438    """Patch the OpenAI client to make it traceable.439 440    Supports:441        - Chat and Responses API's442        - Sync and async OpenAI clients443        - `create` and `parse` methods444        - With and without streaming445        - `with_raw_response` API for accessing HTTP headers446 447    Args:448        client: The client to patch.449        tracing_extra: Extra tracing information.450        chat_name: The run name for the chat completions endpoint.451        completions_name: The run name for the completions endpoint.452 453    Returns:454        The patched client.455 456    Example:457        ```python458        import openai459        from langsmith import wrappers460 461        # Use OpenAI client same as you normally would.462        client = wrappers.wrap_openai(openai.OpenAI())463 464        # Chat API:465        messages = [466            {"role": "system", "content": "You are a helpful assistant."},467            {468                "role": "user",469                "content": "What physics breakthroughs do you predict will happen by 2300?",470            },471        ]472        completion = client.chat.completions.create(473            model="gpt-4o-mini", messages=messages474        )475        print(completion.choices[0].message.content)476 477        # Responses API:478        response = client.responses.create(479            model="gpt-4o-mini",480            messages=messages,481        )482        print(response.output_text)483 484        # With raw response to access headers:485        raw_response = client.chat.completions.with_raw_response.create(486            model="gpt-4o-mini", messages=messages487        )488        print(raw_response.headers)  # Access HTTP headers489        completion = raw_response.parse()  # Get parsed response490        ```491 492    !!! warning "Behavior changed in `langsmith` 0.3.16"493 494        Support for Responses API added.495 496    !!! warning "Behavior changed in `langsmith` 0.3.x"497 498        Support for `with_raw_response` API added.499    """  # noqa: E501500    tracing_extra = tracing_extra or {}501 502    # Extract ls_invocation_params from metadata503    metadata = dict(tracing_extra.get("metadata") or {})504    prepopulated_invocation_params = metadata.pop("ls_invocation_params", {})505 506    # Create new tracing_extra without ls_invocation_params in metadata507    tracing_extra_rest: TracingExtra = {  # type: ignore[assignment]508        k: v for k, v in tracing_extra.items() if k != "metadata"509    }510    if metadata:511        tracing_extra_rest["metadata"] = metadata  # type: ignore[typeddict-item]512 513    ls_provider = "openai"514    try:515        from openai import AsyncAzureOpenAI, AzureOpenAI516 517        if isinstance(client, AzureOpenAI) or isinstance(client, AsyncAzureOpenAI):518            ls_provider = "azure"519            chat_name = "AzureChatOpenAI"520            completions_name = "AzureOpenAI"521    except ImportError:522        pass523 524    # First wrap the create methods - these handle non-streaming cases525    client.chat.completions.create = _get_wrapper(  # type: ignore[method-assign]526        client.chat.completions.create,527        chat_name,528        _reduce_chat,529        tracing_extra=tracing_extra_rest,530        invocation_params_fn=functools.partial(531            _infer_invocation_params,532            "chat",533            ls_provider,534            prepopulated_invocation_params,535            False,536        ),537        process_outputs=_process_chat_completion,538    )539 540    client.completions.create = _get_wrapper(  # type: ignore[method-assign]541        client.completions.create,542        completions_name,543        _reduce_completions,544        tracing_extra=tracing_extra_rest,545        invocation_params_fn=functools.partial(546            _infer_invocation_params,547            "llm",548            ls_provider,549            prepopulated_invocation_params,550            False,551        ),552    )553 554    # Wrap beta.chat.completions.parse if it exists555    if (556        hasattr(client, "beta")557        and hasattr(client.beta, "chat")558        and hasattr(client.beta.chat, "completions")559        and hasattr(client.beta.chat.completions, "parse")560    ):561        client.beta.chat.completions.parse = _get_parse_wrapper(  # type: ignore[method-assign]562            client.beta.chat.completions.parse,  # type: ignore563            chat_name,564            _process_chat_completion,565            tracing_extra=tracing_extra_rest,566            invocation_params_fn=functools.partial(567                _infer_invocation_params,568                "chat",569                ls_provider,570                prepopulated_invocation_params,571                False,572            ),573        )574 575    # Wrap chat.completions.parse if it exists576    if (577        hasattr(client, "chat")578        and hasattr(client.chat, "completions")579        and hasattr(client.chat.completions, "parse")580    ):581        client.chat.completions.parse = _get_parse_wrapper(  # type: ignore[method-assign]582            client.chat.completions.parse,  # type: ignore583            chat_name,584            _process_chat_completion,585            tracing_extra=tracing_extra_rest,586            invocation_params_fn=functools.partial(587                _infer_invocation_params,588                "chat",589                ls_provider,590                prepopulated_invocation_params,591                False,592            ),593        )594 595    # For the responses API: "client.responses.create(**kwargs)"596    if hasattr(client, "responses"):597        if hasattr(client.responses, "create"):598            client.responses.create = _get_wrapper(  # type: ignore[method-assign]599                client.responses.create,600                chat_name,601                _reduce_response_events,602                process_outputs=_process_responses_api_output,603                tracing_extra=tracing_extra_rest,604                invocation_params_fn=functools.partial(605                    _infer_invocation_params,606                    "chat",607                    ls_provider,608                    prepopulated_invocation_params,609                    True,610                ),611            )612        if hasattr(client.responses, "parse"):613            client.responses.parse = _get_parse_wrapper(  # type: ignore[method-assign]614                client.responses.parse,615                chat_name,616                _process_responses_api_output,617                tracing_extra=tracing_extra_rest,618                invocation_params_fn=functools.partial(619                    _infer_invocation_params,620                    "chat",621                    ls_provider,622                    prepopulated_invocation_params,623                    True,624                ),625            )626 627    return client628 629 630def _process_responses_api_output(response: Any) -> dict:631    if response:632        try:633            # Unwrap APIResponse from with_raw_response for tracing634            if hasattr(response, "parse") and callable(response.parse):635                try:636                    response = response.parse()637                except Exception:638                    pass639 640            output = response.model_dump(exclude_none=True, mode="json")641            if usage := output.pop("usage", None):642                output["usage_metadata"] = _create_usage_metadata(643                    usage, output.get("service_tier")644                )645            return output646        except Exception:647            return {"output": response}648    return {}649 
codekingpro/portable-devtools · Team Ai