2023-01-04 15:54:25 +00:00
|
|
|
"""A fake callback handler for testing purposes."""
|
2023-05-11 18:06:39 +00:00
|
|
|
from itertools import chain
|
|
|
|
from typing import Any, Dict, List, Optional
|
|
|
|
from uuid import UUID
|
2023-01-27 01:38:13 +00:00
|
|
|
|
|
|
|
from pydantic import BaseModel
|
2023-01-04 15:54:25 +00:00
|
|
|
|
2023-02-14 23:06:14 +00:00
|
|
|
from langchain.callbacks.base import AsyncCallbackHandler, BaseCallbackHandler
|
2023-07-01 17:39:19 +00:00
|
|
|
from langchain.schema.messages import BaseMessage
|
2023-01-04 15:54:25 +00:00
|
|
|
|
|
|
|
|
2023-02-14 23:06:14 +00:00
|
|
|
class BaseFakeCallbackHandler(BaseModel):
|
|
|
|
"""Base fake callback handler for testing."""
|
2023-01-04 15:54:25 +00:00
|
|
|
|
|
|
|
starts: int = 0
|
|
|
|
ends: int = 0
|
|
|
|
errors: int = 0
|
|
|
|
text: int = 0
|
2023-01-27 01:38:13 +00:00
|
|
|
ignore_llm_: bool = False
|
|
|
|
ignore_chain_: bool = False
|
|
|
|
ignore_agent_: bool = False
|
2023-06-30 21:44:03 +00:00
|
|
|
ignore_retriever_: bool = False
|
2023-05-11 18:06:39 +00:00
|
|
|
ignore_chat_model_: bool = False
|
2023-01-04 15:54:25 +00:00
|
|
|
|
2023-01-28 16:05:20 +00:00
|
|
|
# add finer-grained counters for easier debugging of failing tests
|
|
|
|
chain_starts: int = 0
|
|
|
|
chain_ends: int = 0
|
|
|
|
llm_starts: int = 0
|
|
|
|
llm_ends: int = 0
|
2023-02-14 23:06:14 +00:00
|
|
|
llm_streams: int = 0
|
2023-01-28 16:05:20 +00:00
|
|
|
tool_starts: int = 0
|
|
|
|
tool_ends: int = 0
|
2023-04-30 18:14:09 +00:00
|
|
|
agent_actions: int = 0
|
2023-01-28 16:05:20 +00:00
|
|
|
agent_ends: int = 0
|
2023-05-11 18:06:39 +00:00
|
|
|
chat_model_starts: int = 0
|
2023-06-30 21:44:03 +00:00
|
|
|
retriever_starts: int = 0
|
|
|
|
retriever_ends: int = 0
|
|
|
|
retriever_errors: int = 0
|
2023-01-28 16:05:20 +00:00
|
|
|
|
2023-02-14 23:06:14 +00:00
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
class BaseFakeCallbackHandlerMixin(BaseFakeCallbackHandler):
|
|
|
|
"""Base fake callback handler mixin for testing."""
|
2023-02-14 23:06:14 +00:00
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
def on_llm_start_common(self) -> None:
|
2023-01-28 16:05:20 +00:00
|
|
|
self.llm_starts += 1
|
2023-01-04 15:54:25 +00:00
|
|
|
self.starts += 1
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
def on_llm_end_common(self) -> None:
|
2023-01-28 16:05:20 +00:00
|
|
|
self.llm_ends += 1
|
2023-01-04 15:54:25 +00:00
|
|
|
self.ends += 1
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
def on_llm_error_common(self) -> None:
|
2023-01-04 15:54:25 +00:00
|
|
|
self.errors += 1
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
def on_llm_new_token_common(self) -> None:
|
|
|
|
self.llm_streams += 1
|
|
|
|
|
|
|
|
def on_chain_start_common(self) -> None:
|
2023-06-30 21:44:03 +00:00
|
|
|
("CHAIN START")
|
2023-01-28 16:05:20 +00:00
|
|
|
self.chain_starts += 1
|
2023-01-04 15:54:25 +00:00
|
|
|
self.starts += 1
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
def on_chain_end_common(self) -> None:
|
2023-01-28 16:05:20 +00:00
|
|
|
self.chain_ends += 1
|
2023-01-04 15:54:25 +00:00
|
|
|
self.ends += 1
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
def on_chain_error_common(self) -> None:
|
2023-01-04 15:54:25 +00:00
|
|
|
self.errors += 1
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
def on_tool_start_common(self) -> None:
|
2023-01-28 16:05:20 +00:00
|
|
|
self.tool_starts += 1
|
2023-01-04 15:54:25 +00:00
|
|
|
self.starts += 1
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
def on_tool_end_common(self) -> None:
|
2023-01-28 16:05:20 +00:00
|
|
|
self.tool_ends += 1
|
2023-01-04 15:54:25 +00:00
|
|
|
self.ends += 1
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
def on_tool_error_common(self) -> None:
|
2023-01-04 15:54:25 +00:00
|
|
|
self.errors += 1
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
def on_agent_action_common(self) -> None:
|
2023-05-11 18:06:39 +00:00
|
|
|
print("AGENT ACTION")
|
2023-04-30 18:14:09 +00:00
|
|
|
self.agent_actions += 1
|
|
|
|
self.starts += 1
|
2023-01-04 15:54:25 +00:00
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
def on_agent_finish_common(self) -> None:
|
2023-01-28 16:05:20 +00:00
|
|
|
self.agent_ends += 1
|
2023-01-04 15:54:25 +00:00
|
|
|
self.ends += 1
|
2023-02-14 23:06:14 +00:00
|
|
|
|
2023-05-11 18:06:39 +00:00
|
|
|
def on_chat_model_start_common(self) -> None:
|
|
|
|
print("STARTING CHAT MODEL")
|
|
|
|
self.chat_model_starts += 1
|
|
|
|
self.starts += 1
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
def on_text_common(self) -> None:
|
|
|
|
self.text += 1
|
|
|
|
|
2023-06-30 21:44:03 +00:00
|
|
|
def on_retriever_start_common(self) -> None:
|
|
|
|
self.starts += 1
|
|
|
|
self.retriever_starts += 1
|
|
|
|
|
|
|
|
def on_retriever_end_common(self) -> None:
|
|
|
|
self.ends += 1
|
|
|
|
self.retriever_ends += 1
|
|
|
|
|
|
|
|
def on_retriever_error_common(self) -> None:
|
|
|
|
self.errors += 1
|
|
|
|
self.retriever_errors += 1
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
|
|
|
|
class FakeCallbackHandler(BaseCallbackHandler, BaseFakeCallbackHandlerMixin):
|
|
|
|
"""Fake callback handler for testing."""
|
2023-02-21 06:54:15 +00:00
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
@property
|
|
|
|
def ignore_llm(self) -> bool:
|
|
|
|
"""Whether to ignore LLM callbacks."""
|
|
|
|
return self.ignore_llm_
|
|
|
|
|
|
|
|
@property
|
|
|
|
def ignore_chain(self) -> bool:
|
|
|
|
"""Whether to ignore chain callbacks."""
|
|
|
|
return self.ignore_chain_
|
|
|
|
|
|
|
|
@property
|
|
|
|
def ignore_agent(self) -> bool:
|
|
|
|
"""Whether to ignore agent callbacks."""
|
|
|
|
return self.ignore_agent_
|
2023-02-14 23:06:14 +00:00
|
|
|
|
2023-06-30 21:44:03 +00:00
|
|
|
@property
|
|
|
|
def ignore_retriever(self) -> bool:
|
|
|
|
"""Whether to ignore retriever callbacks."""
|
|
|
|
return self.ignore_retriever_
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
def on_llm_start(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_llm_start_common()
|
|
|
|
|
|
|
|
def on_llm_new_token(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_llm_new_token_common()
|
|
|
|
|
|
|
|
def on_llm_end(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_llm_end_common()
|
|
|
|
|
|
|
|
def on_llm_error(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_llm_error_common()
|
|
|
|
|
|
|
|
def on_chain_start(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_chain_start_common()
|
|
|
|
|
|
|
|
def on_chain_end(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_chain_end_common()
|
|
|
|
|
|
|
|
def on_chain_error(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_chain_error_common()
|
|
|
|
|
|
|
|
def on_tool_start(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_tool_start_common()
|
|
|
|
|
|
|
|
def on_tool_end(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_tool_end_common()
|
|
|
|
|
|
|
|
def on_tool_error(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_tool_error_common()
|
|
|
|
|
|
|
|
def on_agent_action(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_agent_action_common()
|
|
|
|
|
|
|
|
def on_agent_finish(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_agent_finish_common()
|
|
|
|
|
|
|
|
def on_text(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_text_common()
|
|
|
|
|
2023-06-30 21:44:03 +00:00
|
|
|
def on_retriever_start(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_retriever_start_common()
|
|
|
|
|
|
|
|
def on_retriever_end(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_retriever_end_common()
|
|
|
|
|
|
|
|
def on_retriever_error(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
self.on_retriever_error_common()
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
def __deepcopy__(self, memo: dict) -> "FakeCallbackHandler":
|
|
|
|
return self
|
|
|
|
|
|
|
|
|
2023-05-11 18:06:39 +00:00
|
|
|
class FakeCallbackHandlerWithChatStart(FakeCallbackHandler):
|
|
|
|
def on_chat_model_start(
|
|
|
|
self,
|
|
|
|
serialized: Dict[str, Any],
|
|
|
|
messages: List[List[BaseMessage]],
|
|
|
|
*,
|
|
|
|
run_id: UUID,
|
|
|
|
parent_run_id: Optional[UUID] = None,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> Any:
|
|
|
|
assert all(isinstance(m, BaseMessage) for m in chain(*messages))
|
|
|
|
self.on_chat_model_start_common()
|
|
|
|
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
class FakeAsyncCallbackHandler(AsyncCallbackHandler, BaseFakeCallbackHandlerMixin):
|
2023-02-14 23:06:14 +00:00
|
|
|
"""Fake async callback handler for testing."""
|
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
@property
|
|
|
|
def ignore_llm(self) -> bool:
|
|
|
|
"""Whether to ignore LLM callbacks."""
|
|
|
|
return self.ignore_llm_
|
|
|
|
|
|
|
|
@property
|
|
|
|
def ignore_chain(self) -> bool:
|
|
|
|
"""Whether to ignore chain callbacks."""
|
|
|
|
return self.ignore_chain_
|
|
|
|
|
|
|
|
@property
|
|
|
|
def ignore_agent(self) -> bool:
|
|
|
|
"""Whether to ignore agent callbacks."""
|
|
|
|
return self.ignore_agent_
|
|
|
|
|
2023-02-14 23:06:14 +00:00
|
|
|
async def on_llm_start(
|
2023-04-30 18:14:09 +00:00
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
2023-02-14 23:06:14 +00:00
|
|
|
) -> None:
|
2023-04-30 18:14:09 +00:00
|
|
|
self.on_llm_start_common()
|
2023-02-14 23:06:14 +00:00
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
async def on_llm_new_token(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> None:
|
|
|
|
self.on_llm_new_token_common()
|
2023-02-14 23:06:14 +00:00
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
async def on_llm_end(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> None:
|
|
|
|
self.on_llm_end_common()
|
2023-02-14 23:06:14 +00:00
|
|
|
|
|
|
|
async def on_llm_error(
|
2023-04-30 18:14:09 +00:00
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
2023-02-14 23:06:14 +00:00
|
|
|
) -> None:
|
2023-04-30 18:14:09 +00:00
|
|
|
self.on_llm_error_common()
|
2023-02-14 23:06:14 +00:00
|
|
|
|
|
|
|
async def on_chain_start(
|
2023-04-30 18:14:09 +00:00
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
2023-02-14 23:06:14 +00:00
|
|
|
) -> None:
|
2023-04-30 18:14:09 +00:00
|
|
|
self.on_chain_start_common()
|
2023-02-14 23:06:14 +00:00
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
async def on_chain_end(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> None:
|
|
|
|
self.on_chain_end_common()
|
2023-02-14 23:06:14 +00:00
|
|
|
|
|
|
|
async def on_chain_error(
|
2023-04-30 18:14:09 +00:00
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
2023-02-14 23:06:14 +00:00
|
|
|
) -> None:
|
2023-04-30 18:14:09 +00:00
|
|
|
self.on_chain_error_common()
|
2023-02-14 23:06:14 +00:00
|
|
|
|
|
|
|
async def on_tool_start(
|
2023-04-30 18:14:09 +00:00
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
2023-02-14 23:06:14 +00:00
|
|
|
) -> None:
|
2023-04-30 18:14:09 +00:00
|
|
|
self.on_tool_start_common()
|
2023-02-14 23:06:14 +00:00
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
async def on_tool_end(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> None:
|
|
|
|
self.on_tool_end_common()
|
2023-02-14 23:06:14 +00:00
|
|
|
|
|
|
|
async def on_tool_error(
|
2023-04-30 18:14:09 +00:00
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
2023-02-14 23:06:14 +00:00
|
|
|
) -> None:
|
2023-04-30 18:14:09 +00:00
|
|
|
self.on_tool_error_common()
|
2023-02-14 23:06:14 +00:00
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
async def on_agent_action(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> None:
|
|
|
|
self.on_agent_action_common()
|
2023-02-14 23:06:14 +00:00
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
async def on_agent_finish(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> None:
|
|
|
|
self.on_agent_finish_common()
|
2023-02-21 06:54:15 +00:00
|
|
|
|
2023-04-30 18:14:09 +00:00
|
|
|
async def on_text(
|
|
|
|
self,
|
|
|
|
*args: Any,
|
|
|
|
**kwargs: Any,
|
|
|
|
) -> None:
|
|
|
|
self.on_text_common()
|
|
|
|
|
|
|
|
def __deepcopy__(self, memo: dict) -> "FakeAsyncCallbackHandler":
|
|
|
|
return self
|