codekingpro/portable-devtools
114k
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 