mirror of
https://github.com/hwchase17/langchain
synced 2024-10-29 17:07:25 +00:00
6ad360bdef
Please feel free to disregard any changes you disagree with
113 lines
3.2 KiB
Python
113 lines
3.2 KiB
Python
"""A fake callback handler for testing purposes."""
|
|
from typing import Any, Dict, List, Union
|
|
|
|
from pydantic import BaseModel
|
|
|
|
from langchain.callbacks.base import BaseCallbackHandler
|
|
from langchain.schema import AgentAction, AgentFinish, LLMResult
|
|
|
|
|
|
class FakeCallbackHandler(BaseModel, BaseCallbackHandler):
|
|
"""Fake callback handler for testing."""
|
|
|
|
starts: int = 0
|
|
ends: int = 0
|
|
errors: int = 0
|
|
text: int = 0
|
|
ignore_llm_: bool = False
|
|
ignore_chain_: bool = False
|
|
ignore_agent_: bool = False
|
|
always_verbose_: bool = False
|
|
|
|
@property
|
|
def always_verbose(self) -> bool:
|
|
"""Whether to call verbose callbacks even if verbose is False."""
|
|
return self.always_verbose_
|
|
|
|
@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_
|
|
|
|
# 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
|
|
tool_starts: int = 0
|
|
tool_ends: int = 0
|
|
agent_ends: int = 0
|
|
|
|
def on_llm_start(
|
|
self, serialized: Dict[str, Any], prompts: List[str], **kwargs: Any
|
|
) -> None:
|
|
"""Run when LLM starts running."""
|
|
self.llm_starts += 1
|
|
self.starts += 1
|
|
|
|
def on_llm_end(self, response: LLMResult, **kwargs: Any) -> None:
|
|
"""Run when LLM ends running."""
|
|
self.llm_ends += 1
|
|
self.ends += 1
|
|
|
|
def on_llm_error(
|
|
self, error: Union[Exception, KeyboardInterrupt], **kwargs: Any
|
|
) -> None:
|
|
"""Run when LLM errors."""
|
|
self.errors += 1
|
|
|
|
def on_chain_start(
|
|
self, serialized: Dict[str, Any], inputs: Dict[str, Any], **kwargs: Any
|
|
) -> None:
|
|
"""Run when chain starts running."""
|
|
self.chain_starts += 1
|
|
self.starts += 1
|
|
|
|
def on_chain_end(self, outputs: Dict[str, Any], **kwargs: Any) -> None:
|
|
"""Run when chain ends running."""
|
|
self.chain_ends += 1
|
|
self.ends += 1
|
|
|
|
def on_chain_error(
|
|
self, error: Union[Exception, KeyboardInterrupt], **kwargs: Any
|
|
) -> None:
|
|
"""Run when chain errors."""
|
|
self.errors += 1
|
|
|
|
def on_tool_start(
|
|
self, serialized: Dict[str, Any], action: AgentAction, **kwargs: Any
|
|
) -> None:
|
|
"""Run when tool starts running."""
|
|
self.tool_starts += 1
|
|
self.starts += 1
|
|
|
|
def on_tool_end(self, output: str, **kwargs: Any) -> None:
|
|
"""Run when tool ends running."""
|
|
self.tool_ends += 1
|
|
self.ends += 1
|
|
|
|
def on_tool_error(
|
|
self, error: Union[Exception, KeyboardInterrupt], **kwargs: Any
|
|
) -> None:
|
|
"""Run when tool errors."""
|
|
self.errors += 1
|
|
|
|
def on_text(self, text: str, **kwargs: Any) -> None:
|
|
"""Run when agent is ending."""
|
|
self.text += 1
|
|
|
|
def on_agent_finish(self, finish: AgentFinish, **kwargs: Any) -> None:
|
|
"""Run when agent ends running."""
|
|
self.agent_ends += 1
|
|
self.ends += 1
|