codekingpro/portable-devtools
114k
1"""ChatModel wrapper which returns user input as the response.."""2 3from io import StringIO4from typing import Any, Callable, Dict, List, Mapping, Optional5 6import yaml7from langchain_core.callbacks import (8 CallbackManagerForLLMRun,9)10from langchain_core.language_models.chat_models import BaseChatModel11from langchain_core.messages import (12 BaseMessage,13 HumanMessage,14 _message_from_dict,15 messages_to_dict,16)17from langchain_core.outputs import ChatGeneration, ChatResult18from pydantic import Field19 20from langchain_community.llms.utils import enforce_stop_tokens21 22 23def _display_messages(messages: List[BaseMessage]) -> None:24 dict_messages = messages_to_dict(messages)25 for message in dict_messages:26 yaml_string = yaml.dump(27 message,28 default_flow_style=False,29 sort_keys=False,30 allow_unicode=True,31 width=10000,32 line_break=None,33 )34 print("\n", "======= start of message =======", "\n\n") # noqa: T20135 print(yaml_string) # noqa: T20136 print("======= end of message =======", "\n\n") # noqa: T20137 38 39def _collect_yaml_input(40 messages: List[BaseMessage], stop: Optional[List[str]] = None41) -> BaseMessage:42 """Collects and returns user input as a single string."""43 lines = []44 while True:45 line = input()46 if not line.strip():47 break48 if stop and any(seq in line for seq in stop):49 break50 lines.append(line)51 yaml_string = "\n".join(lines)52 53 # Try to parse the input string as YAML54 try:55 message = _message_from_dict(yaml.safe_load(StringIO(yaml_string)))56 if message is None:57 return HumanMessage(content="")58 if stop:59 if isinstance(message.content, str):60 message.content = enforce_stop_tokens(message.content, stop)61 else:62 raise ValueError("Cannot use when output is not a string.")63 return message64 except yaml.YAMLError:65 raise ValueError("Invalid YAML string entered.")66 except ValueError:67 raise ValueError("Invalid message entered.")68 69 70class HumanInputChatModel(BaseChatModel):71 """ChatModel which returns user input as the response."""72 73 input_func: Callable = Field(default_factory=lambda: _collect_yaml_input)74 message_func: Callable = Field(default_factory=lambda: _display_messages)75 separator: str = "\n"76 input_kwargs: Mapping[str, Any] = {}77 message_kwargs: Mapping[str, Any] = {}78 79 @property80 def _identifying_params(self) -> Dict[str, Any]:81 return {82 "input_func": self.input_func.__name__,83 "message_func": self.message_func.__name__,84 }85 86 @property87 def _llm_type(self) -> str:88 """Returns the type of LLM."""89 return "human-input-chat-model"90 91 def _generate(92 self,93 messages: List[BaseMessage],94 stop: Optional[List[str]] = None,95 run_manager: Optional[CallbackManagerForLLMRun] = None,96 **kwargs: Any,97 ) -> ChatResult:98 """99 Displays the messages to the user and returns their input as a response.100 101 Args:102 messages (List[BaseMessage]): The messages to be displayed to the user.103 stop (Optional[List[str]]): A list of stop strings.104 run_manager (Optional[CallbackManagerForLLMRun]): Currently not used.105 106 Returns:107 ChatResult: The user's input as a response.108 """109 self.message_func(messages, **self.message_kwargs)110 user_input = self.input_func(messages, stop=stop, **self.input_kwargs)111 return ChatResult(generations=[ChatGeneration(message=user_input)])112 