Files
langchain/libs/partners/mongodb/tests/integration_tests/test_cache.py
T
Jib d60e93b6ae langchain-mongodb: Standardize mongodb collection/index names in tests (#18755)
## **Description:**
MongoDB integration tests link to a provided Atlas Cluster. We have very
stringent permissions set against the cluster provided. In order to make
it easier to track and isolate the collections each test gets run
against, we've updated the collection names to map the test file name.
i.e. `langchain_{filename}` => `langchain_test_vectorstores`

Fixes integration test results

![image](https://github.com/langchain-ai/langchain/assets/2887713/41f911b9-55f7-4fe4-9134-5514b82009f9)

## **Dependencies:** 
Provided MONGODB_ATLAS_URI

- [x] **Lint and test**: Run `make format`, `make lint` and `make test`
from the root of the package(s) you've modified. See contribution
guidelines for more: https://python.langchain.com/docs/contributing/

cc: @shaneharvey, @blink1073 , @NoahStapp , @caseyclements
2024-03-07 17:16:04 -05:00

155 lines
4.6 KiB
Python

import os
import uuid
from typing import Any, List, Union
import pytest
from langchain_core.caches import BaseCache
from langchain_core.globals import get_llm_cache, set_llm_cache
from langchain_core.load.dump import dumps
from langchain_core.messages import AIMessage, BaseMessage, HumanMessage
from langchain_core.outputs import ChatGeneration, Generation, LLMResult
from langchain_mongodb.cache import MongoDBAtlasSemanticCache, MongoDBCache
from tests.utils import ConsistentFakeEmbeddings, FakeChatModel, FakeLLM
CONN_STRING = os.environ.get("MONGODB_ATLAS_URI")
INDEX_NAME = "langchain-test-index-semantic-cache"
DATABASE = "langchain_test_db"
COLLECTION = "langchain_test_cache"
def random_string() -> str:
return str(uuid.uuid4())
def llm_cache(cls: Any) -> BaseCache:
set_llm_cache(
cls(
embedding=ConsistentFakeEmbeddings(dimensionality=1536),
connection_string=CONN_STRING,
collection_name=COLLECTION,
database_name=DATABASE,
wait_until_ready=True,
)
)
assert get_llm_cache()
return get_llm_cache()
def _execute_test(
prompt: Union[str, List[BaseMessage]],
llm: Union[str, FakeLLM, FakeChatModel],
response: List[Generation],
) -> None:
# Fabricate an LLM String
if not isinstance(llm, str):
params = llm.dict()
params["stop"] = None
llm_string = str(sorted([(k, v) for k, v in params.items()]))
else:
llm_string = llm
# If the prompt is a str then we should pass just the string
dumped_prompt: str = prompt if isinstance(prompt, str) else dumps(prompt)
# Update the cache
get_llm_cache().update(dumped_prompt, llm_string, response)
# Retrieve the cached result through 'generate' call
output: Union[List[Generation], LLMResult, None]
expected_output: Union[List[Generation], LLMResult]
if isinstance(llm, str):
output = get_llm_cache().lookup(dumped_prompt, llm) # type: ignore
expected_output = response
else:
output = llm.generate([prompt]) # type: ignore
expected_output = LLMResult(
generations=[response],
llm_output={},
)
assert output == expected_output # type: ignore
@pytest.mark.parametrize(
"prompt, llm, response",
[
("foo", "bar", [Generation(text="fizz")]),
("foo", FakeLLM(), [Generation(text="fizz")]),
(
[HumanMessage(content="foo")],
FakeChatModel(),
[ChatGeneration(message=AIMessage(content="foo"))],
),
],
ids=[
"plain_cache",
"cache_with_llm",
"cache_with_chat",
],
)
@pytest.mark.parametrize("cacher", [MongoDBCache, MongoDBAtlasSemanticCache])
def test_mongodb_cache(
cacher: Union[MongoDBCache, MongoDBAtlasSemanticCache],
prompt: Union[str, List[BaseMessage]],
llm: Union[str, FakeLLM, FakeChatModel],
response: List[Generation],
) -> None:
llm_cache(cacher)
try:
_execute_test(prompt, llm, response)
finally:
get_llm_cache().clear()
@pytest.mark.parametrize(
"prompts, generations",
[
# Single prompt, single generation
([random_string()], [[random_string()]]),
# Single prompt, multiple generations
([random_string()], [[random_string(), random_string()]]),
# Single prompt, multiple generations
([random_string()], [[random_string(), random_string(), random_string()]]),
# Multiple prompts, multiple generations
(
[random_string(), random_string()],
[[random_string()], [random_string(), random_string()]],
),
],
ids=[
"single_prompt_single_generation",
"single_prompt_two_generations",
"single_prompt_three_generations",
"multiple_prompts_multiple_generations",
],
)
def test_mongodb_atlas_cache_matrix(
prompts: List[str],
generations: List[List[str]],
) -> None:
llm_cache(MongoDBAtlasSemanticCache)
llm = FakeLLM()
# Fabricate an LLM String
params = llm.dict()
params["stop"] = None
llm_string = str(sorted([(k, v) for k, v in params.items()]))
llm_generations = [
[
Generation(text=generation, generation_info=params)
for generation in prompt_i_generations
]
for prompt_i_generations in generations
]
for prompt_i, llm_generations_i in zip(prompts, llm_generations):
_execute_test(prompt_i, llm_string, llm_generations_i)
assert llm.generate(prompts) == LLMResult(
generations=llm_generations, llm_output={}
)
get_llm_cache().clear()