From 10c0b3871c60db216b09b1d001ac5039989b2ced Mon Sep 17 00:00:00 2001 From: Gabriel Luiz Freitas Almeida Date: Fri, 14 Jul 2023 13:10:06 -0300 Subject: [PATCH] =?UTF-8?q?=E2=9C=A8=20feat(conftest.py):=20add=20custom?= =?UTF-8?q?=5Fchain=20fixture=20to=20provide=20a=20custom=20chain=20for=20?= =?UTF-8?q?testing=20purposes?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tests/conftest.py | 119 ++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 119 insertions(+) diff --git a/tests/conftest.py b/tests/conftest.py index f893533ac..ca0bb1dc0 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -116,3 +116,122 @@ def client_fixture(session: Session): # yield TestClient(app) app.dependency_overrides.clear() # + + +@pytest.fixture +def custom_chain(): + return '''from __future__ import annotations + + from typing import Any, Dict, List, Optional + + from pydantic import Extra + + from langchain.schema import BaseLanguageModel, Document + from langchain.callbacks.manager import ( + AsyncCallbackManagerForChainRun, + CallbackManagerForChainRun, + ) + from langchain.chains.base import Chain + from langchain.prompts import StringPromptTemplate + from langflow.interface.custom.base import CustomComponent + + class MyCustomChain(Chain): + """ + An example of a custom chain. + """ + + prompt: StringPromptTemplate + """Prompt object to use.""" + llm: BaseLanguageModel + output_key: str = "text" #: :meta private: + + class Config: + """Configuration for this pydantic object.""" + + extra = Extra.forbid + arbitrary_types_allowed = True + + @property + def input_keys(self) -> List[str]: + """Will be whatever keys the prompt expects. + + :meta private: + """ + return self.prompt.input_variables + + @property + def output_keys(self) -> List[str]: + """Will always return text key. + + :meta private: + """ + return [self.output_key] + + def _call( + self, + inputs: Dict[str, Any], + run_manager: Optional[CallbackManagerForChainRun] = None, + ) -> Dict[str, str]: + # Your custom chain logic goes here + # This is just an example that mimics LLMChain + prompt_value = self.prompt.format_prompt(**inputs) + + # Whenever you call a language model, or another chain, you should pass + # a callback manager to it. This allows the inner run to be tracked by + # any callbacks that are registered on the outer run. + # You can always obtain a callback manager for this by calling + # `run_manager.get_child()` as shown below. + response = self.llm.generate_prompt( + [prompt_value], + callbacks=run_manager.get_child() if run_manager else None, + ) + + # If you want to log something about this run, you can do so by calling + # methods on the `run_manager`, as shown below. This will trigger any + # callbacks that are registered for that event. + if run_manager: + run_manager.on_text("Log something about this run") + + return {self.output_key: response.generations[0][0].text} + + async def _acall( + self, + inputs: Dict[str, Any], + run_manager: Optional[AsyncCallbackManagerForChainRun] = None, + ) -> Dict[str, str]: + # Your custom chain logic goes here + # This is just an example that mimics LLMChain + prompt_value = self.prompt.format_prompt(**inputs) + + # Whenever you call a language model, or another chain, you should pass + # a callback manager to it. This allows the inner run to be tracked by + # any callbacks that are registered on the outer run. + # You can always obtain a callback manager for this by calling + # `run_manager.get_child()` as shown below. + response = await self.llm.agenerate_prompt( + [prompt_value], + callbacks=run_manager.get_child() if run_manager else None, + ) + + # If you want to log something about this run, you can do so by calling + # methods on the `run_manager`, as shown below. This will trigger any + # callbacks that are registered for that event. + if run_manager: + await run_manager.on_text("Log something about this run") + + return {self.output_key: response.generations[0][0].text} + + @property + def _chain_type(self) -> str: + return "my_custom_chain" + + class CustomChain(CustomComponent): + display_name: str = "Custom Chain" + field_config = { + "prompt": {"field_type": "prompt"}, + "llm": {"field_type": "BaseLanguageModel"}, + } + + def build(self, prompt, llm, input: str) -> Document: + chain = MyCustomChain(prompt=prompt, llm=llm) + return chain(input)'''