codekingpro/portable-devtools
114k
1"""Tracer that streams run logs to a stream."""2 3from __future__ import annotations4 5import asyncio6import contextlib7import copy8import threading9from collections import defaultdict10from pprint import pformat11from typing import (12 TYPE_CHECKING,13 Any,14 Literal,15 TypeVar,16 overload,17)18 19import jsonpatch # type: ignore[import-untyped]20from typing_extensions import NotRequired, TypedDict, override21 22from langchain_core.callbacks.base import BaseCallbackManager23from langchain_core.load import dumps24from langchain_core.load.load import load25from langchain_core.outputs import ChatGenerationChunk, GenerationChunk26from langchain_core.runnables import RunnableConfig, ensure_config27from langchain_core.tracers._streaming import _StreamingCallbackHandler28from langchain_core.tracers.base import BaseTracer29from langchain_core.tracers.memory_stream import _MemoryStream30 31if TYPE_CHECKING:32 from collections.abc import AsyncIterator, Iterator, Sequence33 from uuid import UUID34 35 from langchain_core.runnables import Runnable36 from langchain_core.runnables.utils import Input, Output37 from langchain_core.tracers.schemas import Run38 39 40class LogEntry(TypedDict):41 """A single entry in the run log."""42 43 id: str44 """ID of the sub-run."""45 46 name: str47 """Name of the object being run."""48 49 type: str50 """Type of the object being run, eg. prompt, chain, llm, etc."""51 52 tags: list[str]53 """List of tags for the run."""54 55 metadata: dict[str, Any]56 """Key-value pairs of metadata for the run."""57 58 start_time: str59 """ISO-8601 timestamp of when the run started."""60 61 streamed_output_str: list[str]62 """List of LLM tokens streamed by this run, if applicable."""63 64 streamed_output: list[Any]65 """List of output chunks streamed by this run, if available."""66 67 inputs: NotRequired[Any | None]68 """Inputs to this run. Not available currently via `astream_log`."""69 70 final_output: Any | None71 """Final output of this run.72 73 Only available after the run has finished successfully.74 """75 76 end_time: str | None77 """ISO-8601 timestamp of when the run ended.78 79 Only available after the run has finished.80 """81 82 83class RunState(TypedDict):84 """State of the run."""85 86 id: str87 """ID of the run."""88 89 streamed_output: list[Any]90 """List of output chunks streamed by `Runnable.stream()`"""91 92 final_output: Any | None93 """Final output of the run, usually the result of aggregating (`+`) streamed_output.94 95 Updated throughout the run when supported by the `Runnable`.96 """97 98 name: str99 """Name of the object being run."""100 101 type: str102 """Type of the object being run, e.g. prompt, chain, llm, etc."""103 104 # Do we want tags/metadata on the root run? Client kinda knows it in most situations105 # tags: list[str]106 107 logs: dict[str, LogEntry]108 """Map of run names to sub-runs.109 110 If filters were supplied, this list will contain only the runs that matched the111 filters.112 """113 114 115class RunLogPatch:116 """Patch to the run log."""117 118 ops: list[dict[str, Any]]119 """List of `JSONPatch` operations, which describe how to create the run state120 from an empty dict.121 122 This is the minimal representation of the log, designed to be serialized as JSON and123 sent over the wire to reconstruct the log on the other side. Reconstruction of the124 state can be done with any JSONPatch-compliant library, see https://jsonpatch.com125 for more information.126 """127 128 def __init__(self, *ops: dict[str, Any]) -> None:129 """Create a RunLogPatch.130 131 Args:132 *ops: The operations to apply to the state.133 """134 self.ops = list(ops)135 136 def __add__(self, other: RunLogPatch | Any) -> RunLog:137 """Combine two `RunLogPatch` instances.138 139 Args:140 other: The other `RunLogPatch` to combine with.141 142 Raises:143 TypeError: If the other object is not a `RunLogPatch`.144 145 Returns:146 A new `RunLog` representing the combination of the two.147 """148 if type(other) is RunLogPatch:149 ops = self.ops + other.ops150 state = jsonpatch.apply_patch(None, copy.deepcopy(ops))151 return RunLog(*ops, state=state)152 153 msg = f"unsupported operand type(s) for +: '{type(self)}' and '{type(other)}'"154 raise TypeError(msg)155 156 @override157 def __repr__(self) -> str:158 # 1:-1 to get rid of the [] around the list159 return f"RunLogPatch({pformat(self.ops)[1:-1]})"160 161 @override162 def __eq__(self, other: object) -> bool:163 return isinstance(other, RunLogPatch) and self.ops == other.ops164 165 __hash__ = None # type: ignore[assignment]166 167 168class RunLog(RunLogPatch):169 """Run log."""170 171 state: RunState172 """Current state of the log, obtained from applying all ops in sequence."""173 174 def __init__(self, *ops: dict[str, Any], state: RunState) -> None:175 """Create a RunLog.176 177 Args:178 *ops: The operations to apply to the state.179 state: The initial state of the run log.180 """181 super().__init__(*ops)182 self.state = state183 184 def __add__(self, other: RunLogPatch | Any) -> RunLog:185 """Combine two `RunLog` objects.186 187 Args:188 other: The other `RunLog` or `RunLogPatch` to combine with.189 190 Raises:191 TypeError: If the other object is not a `RunLog` or `RunLogPatch`.192 193 Returns:194 A new `RunLog` representing the combination of the two.195 """196 if type(other) is RunLogPatch:197 ops = self.ops + other.ops198 state = jsonpatch.apply_patch(self.state, other.ops)199 return RunLog(*ops, state=state)200 201 msg = f"unsupported operand type(s) for +: '{type(self)}' and '{type(other)}'"202 raise TypeError(msg)203 204 @override205 def __repr__(self) -> str:206 return f"RunLog({pformat(self.state)})"207 208 @override209 def __eq__(self, other: object) -> bool:210 """Check if two `RunLog`s are equal.211 212 Args:213 other: The other `RunLog` to compare to.214 215 Returns:216 `True` if the `RunLog`s are equal, `False` otherwise.217 """218 # First compare that the state is the same219 if not isinstance(other, RunLog):220 return False221 if self.state != other.state:222 return False223 # Then compare that the ops are the same224 return super().__eq__(other)225 226 __hash__ = None227 228 229T = TypeVar("T")230 231 232class LogStreamCallbackHandler(BaseTracer, _StreamingCallbackHandler):233 """Tracer that streams run logs to a stream."""234 235 def __init__(236 self,237 *,238 auto_close: bool = True,239 include_names: Sequence[str] | None = None,240 include_types: Sequence[str] | None = None,241 include_tags: Sequence[str] | None = None,242 exclude_names: Sequence[str] | None = None,243 exclude_types: Sequence[str] | None = None,244 exclude_tags: Sequence[str] | None = None,245 # Schema format is for internal use only.246 _schema_format: Literal["original", "streaming_events"] = "streaming_events",247 ) -> None:248 """A tracer that streams run logs to a stream.249 250 Args:251 auto_close: Whether to close the stream when the root run finishes.252 include_names: Only include runs from `Runnable` objects with matching253 names.254 include_types: Only include runs from `Runnable` objects with matching255 types.256 include_tags: Only include runs from `Runnable` objects with matching tags.257 exclude_names: Exclude runs from `Runnable` objects with matching names.258 exclude_types: Exclude runs from `Runnable` objects with matching types.259 exclude_tags: Exclude runs from `Runnable` objects with matching tags.260 _schema_format: Primarily changes how the inputs and outputs are handled.261 262 **For internal use only. This API will change.**263 264 - `'original'` is the format used by all current tracers. This format is265 slightly inconsistent with respect to inputs and outputs.266 - 'streaming_events' is used for supporting streaming events, for267 internal usage. It will likely change in the future,268 or be deprecated entirely in favor of a dedicated async269 tracer for streaming events.270 271 Raises:272 ValueError: If an invalid schema format is provided (internal use only).273 """274 if _schema_format not in {"original", "streaming_events"}:275 msg = (276 f"Invalid schema format: {_schema_format}. "277 f"Expected one of 'original', 'streaming_events'."278 )279 raise ValueError(msg)280 super().__init__(_schema_format=_schema_format)281 282 self.auto_close = auto_close283 self.include_names = include_names284 self.include_types = include_types285 self.include_tags = include_tags286 self.exclude_names = exclude_names287 self.exclude_types = exclude_types288 self.exclude_tags = exclude_tags289 290 try:291 loop = asyncio.get_event_loop()292 except RuntimeError:293 loop = asyncio.new_event_loop()294 memory_stream = _MemoryStream[RunLogPatch](loop)295 self.lock = threading.Lock()296 self.send_stream = memory_stream.get_send_stream()297 self.receive_stream = memory_stream.get_receive_stream()298 self._key_map_by_run_id: dict[UUID, str] = {}299 self._counter_map_by_name: dict[str, int] = defaultdict(int)300 self.root_id: UUID | None = None301 302 def __aiter__(self) -> AsyncIterator[RunLogPatch]:303 """Iterate over the stream of run logs.304 305 Returns:306 An async iterator over the run log patches.307 """308 return self.receive_stream.__aiter__()309 310 def send(self, *ops: dict[str, Any]) -> bool:311 """Send a patch to the stream, return `False` if the stream is closed.312 313 Args:314 *ops: The operations to send to the stream.315 316 Returns:317 `True` if the patch was sent successfully, `False` if the stream is closed.318 """319 # We will likely want to wrap this in try / except at some point320 # to handle exceptions that might arise at run time.321 # For now we'll let the exception bubble up, and always return322 # True on the happy path.323 self.send_stream.send_nowait(RunLogPatch(*ops))324 return True325 326 async def tap_output_aiter(327 self, run_id: UUID, output: AsyncIterator[T]328 ) -> AsyncIterator[T]:329 """Tap an output async iterator to stream its values to the log.330 331 Args:332 run_id: The ID of the run.333 output: The output async iterator.334 335 Yields:336 The output value.337 """338 async for chunk in output:339 # root run is handled in .astream_log()340 # if we can't find the run silently ignore341 # eg. because this run wasn't included in the log342 if (343 run_id != self.root_id344 and (key := self._key_map_by_run_id.get(run_id))345 and (346 not self.send(347 {348 "op": "add",349 "path": f"/logs/{key}/streamed_output/-",350 "value": chunk,351 }352 )353 )354 ):355 break356 357 yield chunk358 359 def tap_output_iter(self, run_id: UUID, output: Iterator[T]) -> Iterator[T]:360 """Tap an output iterator to stream its values to the log.361 362 Args:363 run_id: The ID of the run.364 output: The output iterator.365 366 Yields:367 The output value.368 """369 for chunk in output:370 # root run is handled in .astream_log()371 # if we can't find the run silently ignore372 # eg. because this run wasn't included in the log373 if (374 run_id != self.root_id375 and (key := self._key_map_by_run_id.get(run_id))376 and (377 not self.send(378 {379 "op": "add",380 "path": f"/logs/{key}/streamed_output/-",381 "value": chunk,382 }383 )384 )385 ):386 break387 388 yield chunk389 390 def include_run(self, run: Run) -> bool:391 """Check if a `Run` should be included in the log.392 393 Args:394 run: The `Run` to check.395 396 Returns:397 `True` if the `Run` should be included, `False` otherwise.398 """399 if run.id == self.root_id:400 return False401 402 run_tags = run.tags or []403 404 if (405 self.include_names is None406 and self.include_types is None407 and self.include_tags is None408 ):409 include = True410 else:411 include = False412 413 if self.include_names is not None:414 include = include or run.name in self.include_names415 if self.include_types is not None:416 include = include or run.run_type in self.include_types417 if self.include_tags is not None:418 include = include or any(tag in self.include_tags for tag in run_tags)419 420 if self.exclude_names is not None:421 include = include and run.name not in self.exclude_names422 if self.exclude_types is not None:423 include = include and run.run_type not in self.exclude_types424 if self.exclude_tags is not None:425 include = include and all(tag not in self.exclude_tags for tag in run_tags)426 427 return include428 429 def _persist_run(self, run: Run) -> None:430 # This is a legacy method only called once for an entire run tree431 # therefore not useful here432 pass433 434 def _on_run_create(self, run: Run) -> None:435 """Start a run."""436 if self.root_id is None:437 self.root_id = run.id438 if not self.send(439 {440 "op": "replace",441 "path": "",442 "value": RunState(443 id=str(run.id),444 streamed_output=[],445 final_output=None,446 logs={},447 name=run.name,448 type=run.run_type,449 ),450 }451 ):452 return453 454 if not self.include_run(run):455 return456 457 # Determine previous index, increment by 1458 with self.lock:459 self._counter_map_by_name[run.name] += 1460 count = self._counter_map_by_name[run.name]461 self._key_map_by_run_id[run.id] = (462 run.name if count == 1 else f"{run.name}:{count}"463 )464 465 entry = LogEntry(466 id=str(run.id),467 name=run.name,468 type=run.run_type,469 tags=run.tags or [],470 metadata=(run.extra or {}).get("metadata", {}),471 start_time=run.start_time.isoformat(timespec="milliseconds"),472 streamed_output=[],473 streamed_output_str=[],474 final_output=None,475 end_time=None,476 )477 478 if self._schema_format == "streaming_events":479 # If using streaming events let's add inputs as well480 entry["inputs"] = _get_standardized_inputs(run, self._schema_format)481 482 # Add the run to the stream483 self.send(484 {485 "op": "add",486 "path": f"/logs/{self._key_map_by_run_id[run.id]}",487 "value": entry,488 }489 )490 491 def _on_run_update(self, run: Run) -> None:492 """Finish a `Run`."""493 try:494 index = self._key_map_by_run_id.get(run.id)495 496 if index is None:497 return498 499 ops = []500 501 if self._schema_format == "streaming_events":502 ops.append(503 {504 "op": "replace",505 "path": f"/logs/{index}/inputs",506 "value": _get_standardized_inputs(run, self._schema_format),507 }508 )509 510 ops.extend(511 [512 # Replace 'inputs' with final inputs513 # This is needed because in many cases the inputs are not514 # known until after the run is finished and the entire515 # input stream has been processed by the runnable.516 {517 "op": "add",518 "path": f"/logs/{index}/final_output",519 # to undo the dumpd done by some runnables / tracer / etc520 "value": _get_standardized_outputs(run, self._schema_format),521 },522 {523 "op": "add",524 "path": f"/logs/{index}/end_time",525 "value": run.end_time.isoformat(timespec="milliseconds")526 if run.end_time is not None527 else None,528 },529 ]530 )531 532 self.send(*ops)533 finally:534 if run.id == self.root_id and self.auto_close:535 self.send_stream.close()536 537 def _on_llm_new_token(538 self,539 run: Run,540 token: str,541 chunk: GenerationChunk | ChatGenerationChunk | None,542 ) -> None:543 """Process new LLM token."""544 index = self._key_map_by_run_id.get(run.id)545 546 if index is None:547 return548 549 self.send(550 {551 "op": "add",552 "path": f"/logs/{index}/streamed_output_str/-",553 "value": token,554 },555 {556 "op": "add",557 "path": f"/logs/{index}/streamed_output/-",558 "value": chunk.message559 if isinstance(chunk, ChatGenerationChunk)560 else token,561 },562 )563 564 565def _get_standardized_inputs(566 run: Run, schema_format: Literal["original", "streaming_events"]567) -> Any:568 """Extract standardized inputs from a `Run`.569 570 Standardizes the inputs based on the type of the runnable used.571 572 Args:573 run: `Run` object574 schema_format: The schema format to use.575 576 Returns:577 Valid inputs are only dict. By conventions, inputs always represented invocation578 using named arguments. `None` means that the input is not yet known!579 """580 if schema_format == "original":581 msg = (582 "Do not assign inputs with original schema drop the key for now."583 "When inputs are added to astream_log they should be added with "584 "standardized schema for streaming events."585 )586 raise NotImplementedError(msg)587 588 inputs = load(run.inputs, allowed_objects="messages")589 590 if run.run_type in {"retriever", "llm", "chat_model"}:591 return inputs592 593 # new style chains594 # These nest an additional 'input' key inside the 'inputs' to make sure595 # the input is always a dict. We need to unpack and use the inner value.596 inputs = inputs["input"]597 # We should try to fix this in Runnables and callbacks/tracers598 # Runnables should be using a None type here not a placeholder599 # dict.600 if inputs == {"input": ""}: # Workaround for Runnables not using None601 # The input is not known, so we don't assign data['input']602 return None603 return inputs604 605 606def _get_standardized_outputs(607 run: Run, schema_format: Literal["original", "streaming_events", "original+chat"]608) -> Any | None:609 """Extract standardized output from a run.610 611 Standardizes the outputs based on the type of the runnable used.612 613 Args:614 run: the run object.615 schema_format: The schema format to use.616 617 Returns:618 An output if returned, otherwise `None`.619 """620 outputs = load(run.outputs, allowed_objects="messages")621 if schema_format == "original":622 if run.run_type == "prompt" and "output" in outputs:623 # These were previously dumped before the tracer.624 # Now we needn't do anything to them.625 return outputs["output"]626 # Return the old schema, without standardizing anything627 return outputs628 629 if run.run_type in {"retriever", "llm", "chat_model"}:630 return outputs631 632 if isinstance(outputs, dict):633 return outputs.get("output", None)634 635 return None636 637 638@overload639def _astream_log_implementation(640 runnable: Runnable[Input, Output],641 value: Any,642 config: RunnableConfig | None = None,643 *,644 stream: LogStreamCallbackHandler,645 diff: Literal[True] = True,646 with_streamed_output_list: bool = True,647 **kwargs: Any,648) -> AsyncIterator[RunLogPatch]: ...649 650 651@overload652def _astream_log_implementation(653 runnable: Runnable[Input, Output],654 value: Any,655 config: RunnableConfig | None = None,656 *,657 stream: LogStreamCallbackHandler,658 diff: Literal[False],659 with_streamed_output_list: bool = True,660 **kwargs: Any,661) -> AsyncIterator[RunLog]: ...662 663 664async def _astream_log_implementation(665 runnable: Runnable[Input, Output],666 value: Any,667 config: RunnableConfig | None = None,668 *,669 stream: LogStreamCallbackHandler,670 diff: bool = True,671 with_streamed_output_list: bool = True,672 **kwargs: Any,673) -> AsyncIterator[RunLogPatch] | AsyncIterator[RunLog]:674 """Implementation of astream_log for a given runnable.675 676 The implementation has been factored out (at least temporarily) as both677 `astream_log` and `astream_events` rely on it.678 679 Args:680 runnable: The runnable to run in streaming mode.681 value: The input to the runnable.682 config: The config to pass to the runnable.683 stream: The stream to send the run logs to.684 diff: Whether to yield run log patches (`True`) or full run logs (`False`).685 with_streamed_output_list: Whether to include a list of all streamed outputs in686 each patch. If `False`, only the final output will be included in the687 patches.688 **kwargs: Additional keyword arguments to pass to the `Runnable`.689 690 Raises:691 ValueError: If the callbacks in the config are of an unexpected type.692 693 Yields:694 The run log patches or states, depending on the value of `diff`.695 """696 # Assign the stream handler to the config697 config = ensure_config(config)698 callbacks = config.get("callbacks")699 if callbacks is None:700 config["callbacks"] = [stream]701 elif isinstance(callbacks, list):702 config["callbacks"] = [*callbacks, stream]703 elif isinstance(callbacks, BaseCallbackManager):704 callbacks = callbacks.copy()705 callbacks.add_handler(stream, inherit=True)706 config["callbacks"] = callbacks707 else:708 msg = (709 f"Unexpected type for callbacks: {callbacks}."710 "Expected None, list or AsyncCallbackManager."711 )712 raise ValueError(msg)713 714 # Call the runnable in streaming mode,715 # add each chunk to the output stream716 async def consume_astream() -> None:717 try:718 prev_final_output: Output | None = None719 final_output: Output | None = None720 721 async for chunk in runnable.astream(value, config, **kwargs):722 prev_final_output = final_output723 if final_output is None:724 final_output = chunk725 else:726 try:727 final_output = final_output + chunk # type: ignore[operator]728 except TypeError:729 prev_final_output = None730 final_output = chunk731 patches: list[dict[str, Any]] = []732 if with_streamed_output_list:733 patches.append(734 {735 "op": "add",736 "path": "/streamed_output/-",737 # chunk cannot be shared between738 # streamed_output and final_output739 # otherwise jsonpatch.apply will740 # modify both741 "value": copy.deepcopy(chunk),742 }743 )744 patches.extend(745 {**op, "path": f"/final_output{op['path']}"}746 for op in jsonpatch.JsonPatch.from_diff(747 prev_final_output, final_output, dumps=dumps748 )749 )750 await stream.send_stream.send(RunLogPatch(*patches))751 finally:752 await stream.send_stream.aclose()753 754 # Start the runnable in a task, so we can start consuming output755 task = asyncio.create_task(consume_astream())756 try:757 # Yield each chunk from the output stream758 if diff:759 async for log in stream:760 yield log761 else:762 state = RunLog(state=None) # type: ignore[arg-type]763 async for log in stream:764 state += log765 yield state766 finally:767 # Wait for the runnable to finish, if not cancelled (eg. by break)768 with contextlib.suppress(asyncio.CancelledError):769 await task770 