codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import asyncio4from collections.abc import AsyncIterator5from typing import Any, Literal, cast6 7from langchain_core.callbacks import AsyncCallbackHandler8from langchain_core.outputs import LLMResult9from typing_extensions import override10 11# TODO: If used by two LLM runs in parallel this won't work as expected12 13 14class AsyncIteratorCallbackHandler(AsyncCallbackHandler):15 """Callback handler that returns an async iterator."""16 17 queue: asyncio.Queue[str]18 19 done: asyncio.Event20 21 @property22 def always_verbose(self) -> bool:23 """Always verbose."""24 return True25 26 def __init__(self) -> None:27 """Instantiate AsyncIteratorCallbackHandler."""28 self.queue = asyncio.Queue()29 self.done = asyncio.Event()30 31 @override32 async def on_llm_start(33 self,34 serialized: dict[str, Any],35 prompts: list[str],36 **kwargs: Any,37 ) -> None:38 # If two calls are made in a row, this resets the state39 self.done.clear()40 41 @override42 async def on_llm_new_token(self, token: str, **kwargs: Any) -> None:43 if token is not None and token != "":44 self.queue.put_nowait(token)45 46 @override47 async def on_llm_end(self, response: LLMResult, **kwargs: Any) -> None:48 self.done.set()49 50 @override51 async def on_llm_error(self, error: BaseException, **kwargs: Any) -> None:52 self.done.set()53 54 # TODO: implement the other methods55 56 async def aiter(self) -> AsyncIterator[str]:57 """Asynchronous iterator that yields tokens."""58 while not self.queue.empty() or not self.done.is_set():59 # Wait for the next token in the queue,60 # but stop waiting if the done event is set61 done, other = await asyncio.wait(62 [63 # NOTE: If you add other tasks here, update the code below,64 # which assumes each set has exactly one task each65 asyncio.ensure_future(self.queue.get()),66 asyncio.ensure_future(self.done.wait()),67 ],68 return_when=asyncio.FIRST_COMPLETED,69 )70 71 # Cancel the other task72 if other:73 other.pop().cancel()74 75 # Extract the value of the first completed task76 token_or_done = cast("str | Literal[True]", done.pop().result())77 78 # If the extracted value is the boolean True, the done event was set79 if token_or_done is True:80 break81 82 # Otherwise, the extracted value is a token, which we yield83 yield token_or_done84 