2022-10-25 02:56:26 +00:00
|
|
|
"""Test functionality related to natbot."""
|
|
|
|
|
2022-11-09 06:17:10 +00:00
|
|
|
from typing import Any, List, Mapping, Optional
|
2022-10-25 02:56:26 +00:00
|
|
|
|
2022-12-13 14:46:01 +00:00
|
|
|
from pydantic import BaseModel
|
|
|
|
|
2022-10-25 02:56:26 +00:00
|
|
|
from langchain.chains.natbot.base import NatBotChain
|
|
|
|
from langchain.llms.base import LLM
|
|
|
|
|
|
|
|
|
2022-12-13 14:46:01 +00:00
|
|
|
class FakeLLM(LLM, BaseModel):
|
2022-10-25 02:56:26 +00:00
|
|
|
"""Fake LLM wrapper for testing purposes."""
|
|
|
|
|
2022-12-15 15:53:32 +00:00
|
|
|
def _call(self, prompt: str, stop: Optional[List[str]] = None) -> str:
|
2022-10-25 02:56:26 +00:00
|
|
|
"""Return `foo` if longer than 10000 words, else `bar`."""
|
|
|
|
if len(prompt) > 10000:
|
|
|
|
return "foo"
|
|
|
|
else:
|
|
|
|
return "bar"
|
|
|
|
|
2022-12-13 14:46:01 +00:00
|
|
|
@property
|
|
|
|
def _llm_type(self) -> str:
|
|
|
|
"""Return type of llm."""
|
|
|
|
return "fake"
|
|
|
|
|
2022-11-09 06:17:10 +00:00
|
|
|
@property
|
|
|
|
def _identifying_params(self) -> Mapping[str, Any]:
|
|
|
|
return {}
|
|
|
|
|
2022-10-25 02:56:26 +00:00
|
|
|
|
|
|
|
def test_proper_inputs() -> None:
|
|
|
|
"""Test that natbot shortens inputs correctly."""
|
|
|
|
nat_bot_chain = NatBotChain(llm=FakeLLM(), objective="testing")
|
|
|
|
url = "foo" * 10000
|
|
|
|
browser_content = "foo" * 10000
|
2022-11-14 02:14:35 +00:00
|
|
|
output = nat_bot_chain.execute(url, browser_content)
|
2022-10-25 02:56:26 +00:00
|
|
|
assert output == "bar"
|
|
|
|
|
|
|
|
|
|
|
|
def test_variable_key_naming() -> None:
|
|
|
|
"""Test that natbot handles variable key naming correctly."""
|
|
|
|
nat_bot_chain = NatBotChain(
|
|
|
|
llm=FakeLLM(),
|
|
|
|
objective="testing",
|
|
|
|
input_url_key="u",
|
|
|
|
input_browser_content_key="b",
|
|
|
|
output_key="c",
|
|
|
|
)
|
2022-11-14 02:14:35 +00:00
|
|
|
output = nat_bot_chain.execute("foo", "foo")
|
2022-10-25 02:56:26 +00:00
|
|
|
assert output == "bar"
|