-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
- [
-
- ]/
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
diff --git a/libs/community/langchain_community/embeddings/__init__.py b/libs/community/langchain_community/embeddings/__init__.py
deleted file mode 100644
index d2d0f3db1e..0000000000
--- a/libs/community/langchain_community/embeddings/__init__.py
+++ /dev/null
@@ -1,452 +0,0 @@
-"""**Embedding models** are wrappers around embedding models
-from different APIs and services.
-
-**Embedding models** can be LLMs or not.
-
-**Class hierarchy:**
-
-.. code-block::
-
- Embeddings --> Embeddings # Examples: OpenAIEmbeddings, HuggingFaceEmbeddings
-"""
-
-import importlib
-import logging
-from typing import TYPE_CHECKING, Any
-
-if TYPE_CHECKING:
- from langchain_community.embeddings.aleph_alpha import (
- AlephAlphaAsymmetricSemanticEmbedding,
- AlephAlphaSymmetricSemanticEmbedding,
- )
- from langchain_community.embeddings.anyscale import (
- AnyscaleEmbeddings,
- )
- from langchain_community.embeddings.ascend import (
- AscendEmbeddings,
- )
- from langchain_community.embeddings.awa import (
- AwaEmbeddings,
- )
- from langchain_community.embeddings.azure_openai import (
- AzureOpenAIEmbeddings,
- )
- from langchain_community.embeddings.baichuan import (
- BaichuanTextEmbeddings,
- )
- from langchain_community.embeddings.baidu_qianfan_endpoint import (
- QianfanEmbeddingsEndpoint,
- )
- from langchain_community.embeddings.bedrock import (
- BedrockEmbeddings,
- )
- from langchain_community.embeddings.bookend import (
- BookendEmbeddings,
- )
- from langchain_community.embeddings.clarifai import (
- ClarifaiEmbeddings,
- )
- from langchain_community.embeddings.clova import (
- ClovaEmbeddings,
- )
- from langchain_community.embeddings.cohere import (
- CohereEmbeddings,
- )
- from langchain_community.embeddings.dashscope import (
- DashScopeEmbeddings,
- )
- from langchain_community.embeddings.databricks import (
- DatabricksEmbeddings,
- )
- from langchain_community.embeddings.deepinfra import (
- DeepInfraEmbeddings,
- )
- from langchain_community.embeddings.edenai import (
- EdenAiEmbeddings,
- )
- from langchain_community.embeddings.elasticsearch import (
- ElasticsearchEmbeddings,
- )
- from langchain_community.embeddings.embaas import (
- EmbaasEmbeddings,
- )
- from langchain_community.embeddings.ernie import (
- ErnieEmbeddings,
- )
- from langchain_community.embeddings.fake import (
- DeterministicFakeEmbedding,
- FakeEmbeddings,
- )
- from langchain_community.embeddings.fastembed import (
- FastEmbedEmbeddings,
- )
- from langchain_community.embeddings.gigachat import (
- GigaChatEmbeddings,
- )
- from langchain_community.embeddings.google_palm import (
- GooglePalmEmbeddings,
- )
- from langchain_community.embeddings.gpt4all import (
- GPT4AllEmbeddings,
- )
- from langchain_community.embeddings.gradient_ai import (
- GradientEmbeddings,
- )
- from langchain_community.embeddings.huggingface import (
- HuggingFaceBgeEmbeddings,
- HuggingFaceEmbeddings,
- HuggingFaceInferenceAPIEmbeddings,
- HuggingFaceInstructEmbeddings,
- )
- from langchain_community.embeddings.huggingface_hub import (
- HuggingFaceHubEmbeddings,
- )
- from langchain_community.embeddings.hunyuan import (
- HunyuanEmbeddings,
- )
- from langchain_community.embeddings.infinity import (
- InfinityEmbeddings,
- )
- from langchain_community.embeddings.infinity_local import (
- InfinityEmbeddingsLocal,
- )
- from langchain_community.embeddings.ipex_llm import IpexLLMBgeEmbeddings
- from langchain_community.embeddings.itrex import (
- QuantizedBgeEmbeddings,
- )
- from langchain_community.embeddings.javelin_ai_gateway import (
- JavelinAIGatewayEmbeddings,
- )
- from langchain_community.embeddings.jina import (
- JinaEmbeddings,
- )
- from langchain_community.embeddings.johnsnowlabs import (
- JohnSnowLabsEmbeddings,
- )
- from langchain_community.embeddings.laser import (
- LaserEmbeddings,
- )
- from langchain_community.embeddings.llamacpp import (
- LlamaCppEmbeddings,
- )
- from langchain_community.embeddings.llamafile import (
- LlamafileEmbeddings,
- )
- from langchain_community.embeddings.llm_rails import (
- LLMRailsEmbeddings,
- )
- from langchain_community.embeddings.localai import (
- LocalAIEmbeddings,
- )
- from langchain_community.embeddings.minimax import (
- MiniMaxEmbeddings,
- )
- from langchain_community.embeddings.mlflow import (
- MlflowCohereEmbeddings,
- MlflowEmbeddings,
- )
- from langchain_community.embeddings.mlflow_gateway import (
- MlflowAIGatewayEmbeddings,
- )
- from langchain_community.embeddings.model2vec import (
- Model2vecEmbeddings,
- )
- from langchain_community.embeddings.modelscope_hub import (
- ModelScopeEmbeddings,
- )
- from langchain_community.embeddings.mosaicml import (
- MosaicMLInstructorEmbeddings,
- )
- from langchain_community.embeddings.naver import (
- ClovaXEmbeddings,
- )
- from langchain_community.embeddings.nemo import (
- NeMoEmbeddings,
- )
- from langchain_community.embeddings.nlpcloud import (
- NLPCloudEmbeddings,
- )
- from langchain_community.embeddings.oci_generative_ai import (
- OCIGenAIEmbeddings,
- )
- from langchain_community.embeddings.octoai_embeddings import (
- OctoAIEmbeddings,
- )
- from langchain_community.embeddings.ollama import (
- OllamaEmbeddings,
- )
- from langchain_community.embeddings.openai import (
- OpenAIEmbeddings,
- )
- from langchain_community.embeddings.openvino import (
- OpenVINOBgeEmbeddings,
- OpenVINOEmbeddings,
- )
- from langchain_community.embeddings.optimum_intel import (
- QuantizedBiEncoderEmbeddings,
- )
- from langchain_community.embeddings.oracleai import (
- OracleEmbeddings,
- )
- from langchain_community.embeddings.ovhcloud import (
- OVHCloudEmbeddings,
- )
- from langchain_community.embeddings.premai import (
- PremAIEmbeddings,
- )
- from langchain_community.embeddings.sagemaker_endpoint import (
- SagemakerEndpointEmbeddings,
- )
- from langchain_community.embeddings.sambanova import (
- SambaStudioEmbeddings,
- )
- from langchain_community.embeddings.self_hosted import (
- SelfHostedEmbeddings,
- )
- from langchain_community.embeddings.self_hosted_hugging_face import (
- SelfHostedHuggingFaceEmbeddings,
- SelfHostedHuggingFaceInstructEmbeddings,
- )
- from langchain_community.embeddings.sentence_transformer import (
- SentenceTransformerEmbeddings,
- )
- from langchain_community.embeddings.solar import (
- SolarEmbeddings,
- )
- from langchain_community.embeddings.spacy_embeddings import (
- SpacyEmbeddings,
- )
- from langchain_community.embeddings.sparkllm import (
- SparkLLMTextEmbeddings,
- )
- from langchain_community.embeddings.tensorflow_hub import (
- TensorflowHubEmbeddings,
- )
- from langchain_community.embeddings.textembed import (
- TextEmbedEmbeddings,
- )
- from langchain_community.embeddings.titan_takeoff import (
- TitanTakeoffEmbed,
- )
- from langchain_community.embeddings.vertexai import (
- VertexAIEmbeddings,
- )
- from langchain_community.embeddings.volcengine import (
- VolcanoEmbeddings,
- )
- from langchain_community.embeddings.voyageai import (
- VoyageEmbeddings,
- )
- from langchain_community.embeddings.xinference import (
- XinferenceEmbeddings,
- )
- from langchain_community.embeddings.yandex import (
- YandexGPTEmbeddings,
- )
- from langchain_community.embeddings.zhipuai import (
- ZhipuAIEmbeddings,
- )
-
-__all__ = [
- "AlephAlphaAsymmetricSemanticEmbedding",
- "AlephAlphaSymmetricSemanticEmbedding",
- "AnyscaleEmbeddings",
- "AscendEmbeddings",
- "AwaEmbeddings",
- "AzureOpenAIEmbeddings",
- "BaichuanTextEmbeddings",
- "BedrockEmbeddings",
- "BookendEmbeddings",
- "ClarifaiEmbeddings",
- "ClovaEmbeddings",
- "ClovaXEmbeddings",
- "CohereEmbeddings",
- "DashScopeEmbeddings",
- "DatabricksEmbeddings",
- "DeepInfraEmbeddings",
- "DeterministicFakeEmbedding",
- "EdenAiEmbeddings",
- "ElasticsearchEmbeddings",
- "EmbaasEmbeddings",
- "ErnieEmbeddings",
- "FakeEmbeddings",
- "FastEmbedEmbeddings",
- "GPT4AllEmbeddings",
- "GigaChatEmbeddings",
- "GooglePalmEmbeddings",
- "GradientEmbeddings",
- "HuggingFaceBgeEmbeddings",
- "HuggingFaceEmbeddings",
- "HuggingFaceHubEmbeddings",
- "HuggingFaceInferenceAPIEmbeddings",
- "HuggingFaceInstructEmbeddings",
- "InfinityEmbeddings",
- "InfinityEmbeddingsLocal",
- "IpexLLMBgeEmbeddings",
- "JavelinAIGatewayEmbeddings",
- "JinaEmbeddings",
- "JohnSnowLabsEmbeddings",
- "LLMRailsEmbeddings",
- "LaserEmbeddings",
- "LlamaCppEmbeddings",
- "LlamafileEmbeddings",
- "LocalAIEmbeddings",
- "MiniMaxEmbeddings",
- "MlflowAIGatewayEmbeddings",
- "MlflowCohereEmbeddings",
- "MlflowEmbeddings",
- "Model2vecEmbeddings",
- "ModelScopeEmbeddings",
- "MosaicMLInstructorEmbeddings",
- "NLPCloudEmbeddings",
- "NeMoEmbeddings",
- "OCIGenAIEmbeddings",
- "OctoAIEmbeddings",
- "OllamaEmbeddings",
- "OpenAIEmbeddings",
- "OpenVINOBgeEmbeddings",
- "OpenVINOEmbeddings",
- "OracleEmbeddings",
- "OVHCloudEmbeddings",
- "PremAIEmbeddings",
- "QianfanEmbeddingsEndpoint",
- "QuantizedBgeEmbeddings",
- "QuantizedBiEncoderEmbeddings",
- "SagemakerEndpointEmbeddings",
- "SambaStudioEmbeddings",
- "SelfHostedEmbeddings",
- "SelfHostedHuggingFaceEmbeddings",
- "SelfHostedHuggingFaceInstructEmbeddings",
- "SentenceTransformerEmbeddings",
- "SolarEmbeddings",
- "SpacyEmbeddings",
- "SparkLLMTextEmbeddings",
- "TensorflowHubEmbeddings",
- "TextEmbedEmbeddings",
- "TitanTakeoffEmbed",
- "VertexAIEmbeddings",
- "VolcanoEmbeddings",
- "VoyageEmbeddings",
- "XinferenceEmbeddings",
- "YandexGPTEmbeddings",
- "ZhipuAIEmbeddings",
- "HunyuanEmbeddings",
-]
-
-_module_lookup = {
- "AlephAlphaAsymmetricSemanticEmbedding": "langchain_community.embeddings.aleph_alpha", # noqa: E501
- "AlephAlphaSymmetricSemanticEmbedding": "langchain_community.embeddings.aleph_alpha", # noqa: E501
- "AnyscaleEmbeddings": "langchain_community.embeddings.anyscale",
- "AwaEmbeddings": "langchain_community.embeddings.awa",
- "AzureOpenAIEmbeddings": "langchain_community.embeddings.azure_openai",
- "BaichuanTextEmbeddings": "langchain_community.embeddings.baichuan",
- "BedrockEmbeddings": "langchain_community.embeddings.bedrock",
- "BookendEmbeddings": "langchain_community.embeddings.bookend",
- "ClarifaiEmbeddings": "langchain_community.embeddings.clarifai",
- "ClovaEmbeddings": "langchain_community.embeddings.clova",
- "ClovaXEmbeddings": "langchain_community.embeddings.naver",
- "CohereEmbeddings": "langchain_community.embeddings.cohere",
- "DashScopeEmbeddings": "langchain_community.embeddings.dashscope",
- "DatabricksEmbeddings": "langchain_community.embeddings.databricks",
- "DeepInfraEmbeddings": "langchain_community.embeddings.deepinfra",
- "DeterministicFakeEmbedding": "langchain_community.embeddings.fake",
- "EdenAiEmbeddings": "langchain_community.embeddings.edenai",
- "ElasticsearchEmbeddings": "langchain_community.embeddings.elasticsearch",
- "EmbaasEmbeddings": "langchain_community.embeddings.embaas",
- "ErnieEmbeddings": "langchain_community.embeddings.ernie",
- "FakeEmbeddings": "langchain_community.embeddings.fake",
- "FastEmbedEmbeddings": "langchain_community.embeddings.fastembed",
- "GPT4AllEmbeddings": "langchain_community.embeddings.gpt4all",
- "GooglePalmEmbeddings": "langchain_community.embeddings.google_palm",
- "GradientEmbeddings": "langchain_community.embeddings.gradient_ai",
- "GigaChatEmbeddings": "langchain_community.embeddings.gigachat",
- "HuggingFaceBgeEmbeddings": "langchain_community.embeddings.huggingface",
- "HuggingFaceEmbeddings": "langchain_community.embeddings.huggingface",
- "HuggingFaceHubEmbeddings": "langchain_community.embeddings.huggingface_hub",
- "HuggingFaceInferenceAPIEmbeddings": "langchain_community.embeddings.huggingface",
- "HuggingFaceInstructEmbeddings": "langchain_community.embeddings.huggingface",
- "InfinityEmbeddings": "langchain_community.embeddings.infinity",
- "InfinityEmbeddingsLocal": "langchain_community.embeddings.infinity_local",
- "IpexLLMBgeEmbeddings": "langchain_community.embeddings.ipex_llm",
- "JavelinAIGatewayEmbeddings": "langchain_community.embeddings.javelin_ai_gateway",
- "JinaEmbeddings": "langchain_community.embeddings.jina",
- "JohnSnowLabsEmbeddings": "langchain_community.embeddings.johnsnowlabs",
- "LLMRailsEmbeddings": "langchain_community.embeddings.llm_rails",
- "LaserEmbeddings": "langchain_community.embeddings.laser",
- "LlamaCppEmbeddings": "langchain_community.embeddings.llamacpp",
- "LlamafileEmbeddings": "langchain_community.embeddings.llamafile",
- "LocalAIEmbeddings": "langchain_community.embeddings.localai",
- "MiniMaxEmbeddings": "langchain_community.embeddings.minimax",
- "MlflowAIGatewayEmbeddings": "langchain_community.embeddings.mlflow_gateway",
- "MlflowCohereEmbeddings": "langchain_community.embeddings.mlflow",
- "MlflowEmbeddings": "langchain_community.embeddings.mlflow",
- "Model2vecEmbeddings": "langchain_community.embeddings.model2vec",
- "ModelScopeEmbeddings": "langchain_community.embeddings.modelscope_hub",
- "MosaicMLInstructorEmbeddings": "langchain_community.embeddings.mosaicml",
- "NLPCloudEmbeddings": "langchain_community.embeddings.nlpcloud",
- "NeMoEmbeddings": "langchain_community.embeddings.nemo",
- "OCIGenAIEmbeddings": "langchain_community.embeddings.oci_generative_ai",
- "OctoAIEmbeddings": "langchain_community.embeddings.octoai_embeddings",
- "OllamaEmbeddings": "langchain_community.embeddings.ollama",
- "OpenAIEmbeddings": "langchain_community.embeddings.openai",
- "OpenVINOEmbeddings": "langchain_community.embeddings.openvino",
- "OpenVINOBgeEmbeddings": "langchain_community.embeddings.openvino",
- "QianfanEmbeddingsEndpoint": "langchain_community.embeddings.baidu_qianfan_endpoint", # noqa: E501
- "QuantizedBgeEmbeddings": "langchain_community.embeddings.itrex",
- "QuantizedBiEncoderEmbeddings": "langchain_community.embeddings.optimum_intel",
- "OracleEmbeddings": "langchain_community.embeddings.oracleai",
- "OVHCloudEmbeddings": "langchain_community.embeddings.ovhcloud",
- "SagemakerEndpointEmbeddings": "langchain_community.embeddings.sagemaker_endpoint",
- "SambaStudioEmbeddings": "langchain_community.embeddings.sambanova",
- "SelfHostedEmbeddings": "langchain_community.embeddings.self_hosted",
- "SelfHostedHuggingFaceEmbeddings": "langchain_community.embeddings.self_hosted_hugging_face", # noqa: E501
- "SelfHostedHuggingFaceInstructEmbeddings": "langchain_community.embeddings.self_hosted_hugging_face", # noqa: E501
- "SentenceTransformerEmbeddings": "langchain_community.embeddings.sentence_transformer", # noqa: E501
- "SolarEmbeddings": "langchain_community.embeddings.solar",
- "SpacyEmbeddings": "langchain_community.embeddings.spacy_embeddings",
- "SparkLLMTextEmbeddings": "langchain_community.embeddings.sparkllm",
- "TensorflowHubEmbeddings": "langchain_community.embeddings.tensorflow_hub",
- "VertexAIEmbeddings": "langchain_community.embeddings.vertexai",
- "VolcanoEmbeddings": "langchain_community.embeddings.volcengine",
- "VoyageEmbeddings": "langchain_community.embeddings.voyageai",
- "XinferenceEmbeddings": "langchain_community.embeddings.xinference",
- "TextEmbedEmbeddings": "langchain_community.embeddings.textembed",
- "TitanTakeoffEmbed": "langchain_community.embeddings.titan_takeoff",
- "PremAIEmbeddings": "langchain_community.embeddings.premai",
- "YandexGPTEmbeddings": "langchain_community.embeddings.yandex",
- "AscendEmbeddings": "langchain_community.embeddings.ascend",
- "ZhipuAIEmbeddings": "langchain_community.embeddings.zhipuai",
- "HunyuanEmbeddings": "langchain_community.embeddings.hunyuan",
-}
-
-
-def __getattr__(name: str) -> Any:
- if name in _module_lookup:
- module = importlib.import_module(_module_lookup[name])
- return getattr(module, name)
- raise AttributeError(f"module {__name__} has no attribute {name}")
-
-
-logger = logging.getLogger(__name__)
-
-
-# TODO: this is in here to maintain backwards compatibility
-class HypotheticalDocumentEmbedder:
- def __init__(self, *args: Any, **kwargs: Any):
- logger.warning(
- "Using a deprecated class. Please use "
- "`from langchain.chains import HypotheticalDocumentEmbedder` instead"
- )
- from langchain.chains.hyde.base import HypotheticalDocumentEmbedder as H
-
- return H(*args, **kwargs) # type: ignore[return-value]
-
- @classmethod
- def from_llm(cls, *args: Any, **kwargs: Any) -> Any:
- logger.warning(
- "Using a deprecated class. Please use "
- "`from langchain.chains import HypotheticalDocumentEmbedder` instead"
- )
- from langchain.chains.hyde.base import HypotheticalDocumentEmbedder as H
-
- return H.from_llm(*args, **kwargs)
diff --git a/libs/community/langchain_community/embeddings/aleph_alpha.py b/libs/community/langchain_community/embeddings/aleph_alpha.py
deleted file mode 100644
index 96426fdac8..0000000000
--- a/libs/community/langchain_community/embeddings/aleph_alpha.py
+++ /dev/null
@@ -1,256 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import BaseModel, model_validator
-
-
-class AlephAlphaAsymmetricSemanticEmbedding(BaseModel, Embeddings):
- """Aleph Alpha's asymmetric semantic embedding.
-
- AA provides you with an endpoint to embed a document and a query.
- The models were optimized to make the embeddings of documents and
- the query for a document as similar as possible.
- To learn more, check out: https://docs.aleph-alpha.com/docs/tasks/semantic_embed/
-
- Example:
- .. code-block:: python
- from aleph_alpha import AlephAlphaAsymmetricSemanticEmbedding
-
- embeddings = AlephAlphaAsymmetricSemanticEmbedding(
- normalize=True, compress_to_size=128
- )
-
- document = "This is a content of the document"
- query = "What is the content of the document?"
-
- doc_result = embeddings.embed_documents([document])
- query_result = embeddings.embed_query(query)
-
- """
-
- client: Any #: :meta private:
-
- # Embedding params
- model: str = "luminous-base"
- """Model name to use."""
- compress_to_size: Optional[int] = None
- """Should the returned embeddings come back as an original 5120-dim vector,
- or should it be compressed to 128-dim."""
- normalize: bool = False
- """Should returned embeddings be normalized"""
- contextual_control_threshold: Optional[int] = None
- """Attention control parameters only apply to those tokens that have
- explicitly been set in the request."""
- control_log_additive: bool = True
- """Apply controls on prompt items by adding the log(control_factor)
- to attention scores."""
-
- # Client params
- aleph_alpha_api_key: Optional[str] = None
- """API key for Aleph Alpha API."""
- host: str = "https://api.aleph-alpha.com"
- """The hostname of the API host.
- The default one is "https://api.aleph-alpha.com")"""
- hosting: Optional[str] = None
- """Determines in which datacenters the request may be processed.
- You can either set the parameter to "aleph-alpha" or omit it (defaulting to None).
- Not setting this value, or setting it to None, gives us maximal flexibility
- in processing your request in our
- own datacenters and on servers hosted with other providers.
- Choose this option for maximal availability.
- Setting it to "aleph-alpha" allows us to only process the request
- in our own datacenters.
- Choose this option for maximal data privacy."""
- request_timeout_seconds: int = 305
- """Client timeout that will be set for HTTP requests in the
- `requests` library's API calls.
- Server will close all requests after 300 seconds with an internal server error."""
- total_retries: int = 8
- """The number of retries made in case requests fail with certain retryable
- status codes. If the last
- retry fails a corresponding exception is raised. Note, that between retries
- an exponential backoff
- is applied, starting with 0.5 s after the first retry and doubling for each
- retry made. So with the
- default setting of 8 retries a total wait time of 63.5 s is added between
- the retries."""
- nice: bool = False
- """Setting this to True, will signal to the API that you intend to be
- nice to other users
- by de-prioritizing your request below concurrent ones."""
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and python package exists in environment."""
- aleph_alpha_api_key = get_from_dict_or_env(
- values, "aleph_alpha_api_key", "ALEPH_ALPHA_API_KEY"
- )
- try:
- from aleph_alpha_client import Client
-
- values["client"] = Client(
- token=aleph_alpha_api_key,
- host=values["host"],
- hosting=values["hosting"],
- request_timeout_seconds=values["request_timeout_seconds"],
- total_retries=values["total_retries"],
- nice=values["nice"],
- )
- except ImportError:
- raise ImportError(
- "Could not import aleph_alpha_client python package. "
- "Please install it with `pip install aleph_alpha_client`."
- )
-
- return values
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Call out to Aleph Alpha's asymmetric Document endpoint.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- try:
- from aleph_alpha_client import (
- Prompt,
- SemanticEmbeddingRequest,
- SemanticRepresentation,
- )
- except ImportError:
- raise ImportError(
- "Could not import aleph_alpha_client python package. "
- "Please install it with `pip install aleph_alpha_client`."
- )
- document_embeddings = []
-
- for text in texts:
- document_params = {
- "prompt": Prompt.from_text(text),
- "representation": SemanticRepresentation.Document,
- "compress_to_size": self.compress_to_size,
- "normalize": self.normalize,
- "contextual_control_threshold": self.contextual_control_threshold,
- "control_log_additive": self.control_log_additive,
- }
-
- document_request = SemanticEmbeddingRequest(**document_params)
- document_response = self.client.semantic_embed(
- request=document_request, model=self.model
- )
-
- document_embeddings.append(document_response.embedding)
-
- return document_embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to Aleph Alpha's asymmetric, query embedding endpoint
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- try:
- from aleph_alpha_client import (
- Prompt,
- SemanticEmbeddingRequest,
- SemanticRepresentation,
- )
- except ImportError:
- raise ImportError(
- "Could not import aleph_alpha_client python package. "
- "Please install it with `pip install aleph_alpha_client`."
- )
- symmetric_params = {
- "prompt": Prompt.from_text(text),
- "representation": SemanticRepresentation.Query,
- "compress_to_size": self.compress_to_size,
- "normalize": self.normalize,
- "contextual_control_threshold": self.contextual_control_threshold,
- "control_log_additive": self.control_log_additive,
- }
-
- symmetric_request = SemanticEmbeddingRequest(**symmetric_params)
- symmetric_response = self.client.semantic_embed(
- request=symmetric_request, model=self.model
- )
-
- return symmetric_response.embedding
-
-
-class AlephAlphaSymmetricSemanticEmbedding(AlephAlphaAsymmetricSemanticEmbedding):
- """Symmetric version of the Aleph Alpha's semantic embeddings.
-
- The main difference is that here, both the documents and
- queries are embedded with a SemanticRepresentation.Symmetric
- Example:
- .. code-block:: python
-
- from aleph_alpha import AlephAlphaSymmetricSemanticEmbedding
-
- embeddings = AlephAlphaAsymmetricSemanticEmbedding(
- normalize=True, compress_to_size=128
- )
- text = "This is a test text"
-
- doc_result = embeddings.embed_documents([text])
- query_result = embeddings.embed_query(text)
- """
-
- def _embed(self, text: str) -> List[float]:
- try:
- from aleph_alpha_client import (
- Prompt,
- SemanticEmbeddingRequest,
- SemanticRepresentation,
- )
- except ImportError:
- raise ImportError(
- "Could not import aleph_alpha_client python package. "
- "Please install it with `pip install aleph_alpha_client`."
- )
- query_params = {
- "prompt": Prompt.from_text(text),
- "representation": SemanticRepresentation.Symmetric,
- "compress_to_size": self.compress_to_size,
- "normalize": self.normalize,
- "contextual_control_threshold": self.contextual_control_threshold,
- "control_log_additive": self.control_log_additive,
- }
-
- query_request = SemanticEmbeddingRequest(**query_params)
- query_response = self.client.semantic_embed(
- request=query_request, model=self.model
- )
-
- return query_response.embedding
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Call out to Aleph Alpha's Document endpoint.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- document_embeddings = []
-
- for text in texts:
- document_embeddings.append(self._embed(text))
- return document_embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to Aleph Alpha's asymmetric, query embedding endpoint
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self._embed(text)
diff --git a/libs/community/langchain_community/embeddings/anyscale.py b/libs/community/langchain_community/embeddings/anyscale.py
deleted file mode 100644
index ffa33fa497..0000000000
--- a/libs/community/langchain_community/embeddings/anyscale.py
+++ /dev/null
@@ -1,76 +0,0 @@
-"""Anyscale embeddings wrapper."""
-
-from __future__ import annotations
-
-from typing import Dict, Optional
-
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import Field, SecretStr
-
-from langchain_community.embeddings.openai import OpenAIEmbeddings
-from langchain_community.utils.openai import is_openai_v1
-
-DEFAULT_API_BASE = "https://api.endpoints.anyscale.com/v1"
-DEFAULT_MODEL = "thenlper/gte-large"
-
-
-class AnyscaleEmbeddings(OpenAIEmbeddings):
- """`Anyscale` Embeddings API."""
-
- anyscale_api_key: Optional[SecretStr] = Field(default=None)
- """AnyScale Endpoints API keys."""
- model: str = Field(default=DEFAULT_MODEL)
- """Model name to use."""
- anyscale_api_base: str = Field(default=DEFAULT_API_BASE)
- """Base URL path for API requests."""
- tiktoken_enabled: bool = False
- """Set this to False for non-OpenAI implementations of the embeddings API"""
- embedding_ctx_length: int = 500
- """The maximum number of tokens to embed at once."""
-
- @property
- def lc_secrets(self) -> Dict[str, str]:
- return {
- "anyscale_api_key": "ANYSCALE_API_KEY",
- }
-
- @pre_init
- def validate_environment(cls, values: dict) -> dict:
- """Validate that api key and python package exists in environment."""
- values["anyscale_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(
- values,
- "anyscale_api_key",
- "ANYSCALE_API_KEY",
- )
- )
- values["anyscale_api_base"] = get_from_dict_or_env(
- values,
- "anyscale_api_base",
- "ANYSCALE_API_BASE",
- default=DEFAULT_API_BASE,
- )
- try:
- import openai
-
- except ImportError:
- raise ImportError(
- "Could not import openai python package. "
- "Please install it with `pip install openai`."
- )
- if is_openai_v1():
- # For backwards compatibility.
- client_params = {
- "api_key": values["anyscale_api_key"].get_secret_value(),
- "base_url": values["anyscale_api_base"],
- }
- values["client"] = openai.OpenAI(**client_params).embeddings
- else:
- values["openai_api_base"] = values["anyscale_api_base"]
- values["openai_api_key"] = values["anyscale_api_key"].get_secret_value()
- values["client"] = openai.Embedding
- return values
-
- @property
- def _llm_type(self) -> str:
- return "anyscale-embedding"
diff --git a/libs/community/langchain_community/embeddings/ascend.py b/libs/community/langchain_community/embeddings/ascend.py
deleted file mode 100644
index 940b84bbfc..0000000000
--- a/libs/community/langchain_community/embeddings/ascend.py
+++ /dev/null
@@ -1,137 +0,0 @@
-import os
-from typing import Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict, model_validator
-
-
-class AscendEmbeddings(Embeddings, BaseModel):
- """
- Ascend NPU accelerate Embedding model
-
- Please ensure that you have installed CANN and torch_npu.
-
- Example:
-
- from langchain_community.embeddings import AscendEmbeddings
- model = AscendEmbeddings(model_path=,
- device_id=0,
- query_instruction="Represent this sentence for searching relevant passages: "
- )
- """
-
- """model path"""
- model_path: str
- """Ascend NPU device id."""
- device_id: int = 0
- """Unstruntion to used for embedding query."""
- query_instruction: str = ""
- """Unstruntion to used for embedding document."""
- document_instruction: str = ""
- use_fp16: bool = True
- pooling_method: Optional[str] = "cls"
- batch_size: int = 32
- model: Any
- tokenizer: Any
-
- model_config = ConfigDict(protected_namespaces=())
-
- def __init__(self, *args: Any, **kwargs: Any) -> None:
- super().__init__(*args, **kwargs)
- try:
- from transformers import AutoModel, AutoTokenizer
- except ImportError as e:
- raise ImportError(
- "Unable to import transformers, please install with "
- "`pip install -U transformers`."
- ) from e
- try:
- self.model = AutoModel.from_pretrained(self.model_path).npu().eval()
- self.tokenizer = AutoTokenizer.from_pretrained(self.model_path)
- except Exception as e:
- raise Exception(
- f"Failed to load model [self.model_path], due to following error:{e}"
- )
-
- if self.use_fp16:
- self.model.half()
- self.encode([f"warmup {i} times" for i in range(10)])
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- if "model_path" not in values:
- raise ValueError("model_path is required")
- if not os.access(values["model_path"], os.F_OK):
- raise FileNotFoundError(
- f"Unable to find valid model path in [{values['model_path']}]"
- )
- try:
- import torch_npu
- except ImportError:
- raise ModuleNotFoundError("torch_npu not found, please install torch_npu")
- except Exception as e:
- raise e
- try:
- torch_npu.npu.set_device(values["device_id"])
- except Exception as e:
- raise Exception(f"set device failed due to {e}")
- return values
-
- def encode(self, sentences: Any) -> Any:
- inputs = self.tokenizer(
- sentences,
- padding=True,
- truncation=True,
- return_tensors="pt",
- max_length=512,
- )
- try:
- import torch
- except ImportError as e:
- raise ImportError(
- "Unable to import torch, please install with `pip install -U torch`."
- ) from e
- last_hidden_state = self.model(
- inputs.input_ids.npu(), inputs.attention_mask.npu(), return_dict=True
- ).last_hidden_state
- tmp = self.pooling(last_hidden_state, inputs["attention_mask"].npu())
- embeddings = torch.nn.functional.normalize(tmp, dim=-1)
- return embeddings.cpu().detach().numpy()
-
- def pooling(self, last_hidden_state: Any, attention_mask: Any = None) -> Any:
- try:
- import torch
- except ImportError as e:
- raise ImportError(
- "Unable to import torch, please install with `pip install -U torch`."
- ) from e
- if self.pooling_method == "cls":
- return last_hidden_state[:, 0]
- elif self.pooling_method == "mean":
- s = torch.sum(
- last_hidden_state * attention_mask.unsqueeze(-1).float(), dim=-1
- )
- d = attention_mask.sum(dim=1, keepdim=True).float()
- return s / d
- else:
- raise NotImplementedError(
- f"Pooling method [{self.pooling_method}] not implemented"
- )
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- try:
- import numpy as np
- except ImportError as e:
- raise ImportError(
- "Unable to import numpy, please install with `pip install -U numpy`."
- ) from e
- embedding_list = []
- for i in range(0, len(texts), self.batch_size):
- texts_ = texts[i : i + self.batch_size]
- emb = self.encode([self.document_instruction + text for text in texts_])
- embedding_list.append(emb)
- return np.concatenate(embedding_list)
-
- def embed_query(self, text: str) -> List[float]:
- return self.encode([self.query_instruction + text])[0]
diff --git a/libs/community/langchain_community/embeddings/awa.py b/libs/community/langchain_community/embeddings/awa.py
deleted file mode 100644
index 27cb422423..0000000000
--- a/libs/community/langchain_community/embeddings/awa.py
+++ /dev/null
@@ -1,64 +0,0 @@
-from typing import Any, Dict, List
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, model_validator
-
-
-class AwaEmbeddings(BaseModel, Embeddings):
- """Embedding documents and queries with Awa DB.
-
- Attributes:
- client: The AwaEmbedding client.
- model: The name of the model used for embedding.
- Default is "all-mpnet-base-v2".
- """
-
- client: Any #: :meta private:
- model: str = "all-mpnet-base-v2"
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that awadb library is installed."""
-
- try:
- from awadb import AwaEmbedding
- except ImportError as exc:
- raise ImportError(
- "Could not import awadb library. "
- "Please install it with `pip install awadb`"
- ) from exc
- values["client"] = AwaEmbedding()
- return values
-
- def set_model(self, model_name: str) -> None:
- """Set the model used for embedding.
- The default model used is all-mpnet-base-v2
-
- Args:
- model_name: A string which represents the name of model.
- """
- self.model = model_name
- self.client.model_name = model_name
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed a list of documents using AwaEmbedding.
-
- Args:
- texts: The list of texts need to be embedded
-
- Returns:
- List of embeddings, one for each text.
- """
- return self.client.EmbeddingBatch(texts)
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using AwaEmbedding.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self.client.Embedding(text)
diff --git a/libs/community/langchain_community/embeddings/azure_openai.py b/libs/community/langchain_community/embeddings/azure_openai.py
deleted file mode 100644
index 00a2327d2c..0000000000
--- a/libs/community/langchain_community/embeddings/azure_openai.py
+++ /dev/null
@@ -1,187 +0,0 @@
-"""Azure OpenAI embeddings wrapper."""
-
-from __future__ import annotations
-
-import os
-import warnings
-from typing import Any, Awaitable, Callable, Dict, Optional, Union
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import Field, model_validator
-from typing_extensions import Self
-
-from langchain_community.embeddings.openai import OpenAIEmbeddings
-from langchain_community.utils.openai import is_openai_v1
-
-
-@deprecated(
- since="0.0.9",
- removal="1.0",
- alternative_import="langchain_openai.AzureOpenAIEmbeddings",
-)
-class AzureOpenAIEmbeddings(OpenAIEmbeddings):
- """`Azure OpenAI` Embeddings API."""
-
- azure_endpoint: Union[str, None] = None
- """Your Azure endpoint, including the resource.
-
- Automatically inferred from env var `AZURE_OPENAI_ENDPOINT` if not provided.
-
- Example: `https://example-resource.azure.openai.com/`
- """
- deployment: Optional[str] = Field(default=None, alias="azure_deployment")
- """A model deployment.
-
- If given sets the base client URL to include `/deployments/{azure_deployment}`.
- Note: this means you won't be able to use non-deployment endpoints.
- """
- openai_api_key: Union[str, None] = Field(default=None, alias="api_key")
- """Automatically inferred from env var `AZURE_OPENAI_API_KEY` if not provided."""
- azure_ad_token: Union[str, None] = None
- """Your Azure Active Directory token.
-
- Automatically inferred from env var `AZURE_OPENAI_AD_TOKEN` if not provided.
-
- For more:
- https://www.microsoft.com/en-us/security/business/identity-access/microsoft-entra-id.
- """
- azure_ad_token_provider: Union[Callable[[], str], None] = None
- """A function that returns an Azure Active Directory token.
-
- Will be invoked on every sync request. For async requests,
- will be invoked if `azure_ad_async_token_provider` is not provided.
- """
- azure_ad_async_token_provider: Union[Callable[[], Awaitable[str]], None] = None
- """A function that returns an Azure Active Directory token.
-
- Will be invoked on every async request.
- """
- openai_api_version: Optional[str] = Field(default=None, alias="api_version")
- """Automatically inferred from env var `OPENAI_API_VERSION` if not provided."""
- validate_base_url: bool = True
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and python package exists in environment."""
- # Check OPENAI_KEY for backwards compatibility.
- # TODO: Remove OPENAI_API_KEY support to avoid possible conflict when using
- # other forms of azure credentials.
- values["openai_api_key"] = (
- values.get("openai_api_key")
- or os.getenv("AZURE_OPENAI_API_KEY")
- or os.getenv("OPENAI_API_KEY")
- )
- values["openai_api_base"] = values.get("openai_api_base") or os.getenv(
- "OPENAI_API_BASE"
- )
- values["openai_api_version"] = values.get("openai_api_version") or os.getenv(
- "OPENAI_API_VERSION", default="2023-05-15"
- )
- values["openai_api_type"] = get_from_dict_or_env(
- values, "openai_api_type", "OPENAI_API_TYPE", default="azure"
- )
- values["openai_organization"] = (
- values.get("openai_organization")
- or os.getenv("OPENAI_ORG_ID")
- or os.getenv("OPENAI_ORGANIZATION")
- )
- values["openai_proxy"] = get_from_dict_or_env(
- values,
- "openai_proxy",
- "OPENAI_PROXY",
- default="",
- )
- values["azure_endpoint"] = values.get("azure_endpoint") or os.getenv(
- "AZURE_OPENAI_ENDPOINT"
- )
- values["azure_ad_token"] = values.get("azure_ad_token") or os.getenv(
- "AZURE_OPENAI_AD_TOKEN"
- )
- # Azure OpenAI embedding models allow a maximum of 2048 texts
- # at a time in each batch
- # See: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#embeddings
- values["chunk_size"] = min(values["chunk_size"], 2048)
- try:
- import openai # noqa: F401
- except ImportError:
- raise ImportError(
- "Could not import openai python package. "
- "Please install it with `pip install openai`."
- )
- if is_openai_v1():
- # For backwards compatibility. Before openai v1, no distinction was made
- # between azure_endpoint and base_url (openai_api_base).
- openai_api_base = values["openai_api_base"]
- if openai_api_base and values["validate_base_url"]:
- if "/openai" not in openai_api_base:
- values["openai_api_base"] += "/openai"
- warnings.warn(
- "As of openai>=1.0.0, Azure endpoints should be specified via "
- f"the `azure_endpoint` param not `openai_api_base` "
- f"(or alias `base_url`). Updating `openai_api_base` from "
- f"{openai_api_base} to {values['openai_api_base']}."
- )
- if values["deployment"]:
- warnings.warn(
- "As of openai>=1.0.0, if `deployment` (or alias "
- "`azure_deployment`) is specified then "
- "`openai_api_base` (or alias `base_url`) should not be. "
- "Instead use `deployment` (or alias `azure_deployment`) "
- "and `azure_endpoint`."
- )
- if values["deployment"] not in values["openai_api_base"]:
- warnings.warn(
- "As of openai>=1.0.0, if `openai_api_base` "
- "(or alias `base_url`) is specified it is expected to be "
- "of the form "
- "https://example-resource.azure.openai.com/openai/deployments/example-deployment. " # noqa: E501
- f"Updating {openai_api_base} to "
- f"{values['openai_api_base']}."
- )
- values["openai_api_base"] += (
- "/deployments/" + values["deployment"]
- )
- values["deployment"] = None
- return values
-
- @model_validator(mode="after")
- def post_init_validator(self) -> Self:
- """Validate that the base url is set."""
- import openai
-
- if is_openai_v1():
- client_params = {
- "api_version": self.openai_api_version,
- "azure_endpoint": self.azure_endpoint,
- "azure_deployment": self.deployment,
- "api_key": self.openai_api_key,
- "azure_ad_token": self.azure_ad_token,
- "azure_ad_token_provider": self.azure_ad_token_provider,
- "organization": self.openai_organization,
- "base_url": self.openai_api_base,
- "timeout": self.request_timeout,
- "max_retries": self.max_retries,
- "default_headers": {
- **(self.default_headers or {}),
- "User-Agent": "langchain-comm-python-azure-openai",
- },
- "default_query": self.default_query,
- "http_client": self.http_client,
- }
- self.client = openai.AzureOpenAI(**client_params).embeddings
-
- if self.azure_ad_async_token_provider:
- client_params["azure_ad_token_provider"] = (
- self.azure_ad_async_token_provider
- )
-
- self.async_client = openai.AsyncAzureOpenAI(**client_params).embeddings
- else:
- self.client = openai.Embedding
- return self
-
- @property
- def _llm_type(self) -> str:
- return "azure-openai-chat"
diff --git a/libs/community/langchain_community/embeddings/baichuan.py b/libs/community/langchain_community/embeddings/baichuan.py
deleted file mode 100644
index c12aaa44f1..0000000000
--- a/libs/community/langchain_community/embeddings/baichuan.py
+++ /dev/null
@@ -1,150 +0,0 @@
-from typing import Any, List, Optional
-
-import requests
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import (
- secret_from_env,
-)
-from pydantic import (
- BaseModel,
- ConfigDict,
- Field,
- SecretStr,
- model_validator,
-)
-from requests import RequestException
-from typing_extensions import Self
-
-BAICHUAN_API_URL: str = "https://api.baichuan-ai.com/v1/embeddings"
-
-# BaichuanTextEmbeddings is an embedding model provided by Baichuan Inc. (https://www.baichuan-ai.com/home).
-# As of today (Jan 25th, 2024) BaichuanTextEmbeddings ranks #1 in C-MTEB
-# (Chinese Multi-Task Embedding Benchmark) leaderboard.
-# Leaderboard (Under Overall -> Chinese section): https://huggingface.co/spaces/mteb/leaderboard
-
-# Official Website: https://platform.baichuan-ai.com/docs/text-Embedding
-# An API-key is required to use this embedding model. You can get one by registering
-# at https://platform.baichuan-ai.com/docs/text-Embedding.
-# BaichuanTextEmbeddings support 512 token window and produces vectors with
-# 1024 dimensions.
-
-
-# NOTE!! BaichuanTextEmbeddings only supports Chinese text embedding.
-# Multi-language support is coming soon.
-class BaichuanTextEmbeddings(BaseModel, Embeddings):
- """Baichuan Text Embedding models.
-
- Setup:
- To use, you should set the environment variable ``BAICHUAN_API_KEY`` to
- your API key or pass it as a named parameter to the constructor.
-
- .. code-block:: bash
-
- export BAICHUAN_API_KEY="your-api-key"
-
- Instantiate:
- .. code-block:: python
-
- from langchain_community.embeddings import BaichuanTextEmbeddings
-
- embeddings = BaichuanTextEmbeddings()
-
- Embed:
- .. code-block:: python
-
- # embed the documents
- vectors = embeddings.embed_documents([text1, text2, ...])
-
- # embed the query
- vectors = embeddings.embed_query(text)
- """ # noqa: E501
-
- session: Any = None #: :meta private:
- model_name: str = Field(default="Baichuan-Text-Embedding", alias="model")
- """The model used to embed the documents."""
- baichuan_api_key: SecretStr = Field(
- alias="api_key",
- default_factory=secret_from_env(["BAICHUAN_API_KEY", "BAICHUAN_AUTH_TOKEN"]),
- )
- """Automatically inferred from env var `BAICHUAN_API_KEY` if not provided."""
- chunk_size: int = 16
- """Chunk size when multiple texts are input"""
-
- model_config = ConfigDict(populate_by_name=True, protected_namespaces=())
-
- @model_validator(mode="after")
- def validate_environment(self) -> Self:
- """Validate that auth token exists in environment."""
- session = requests.Session()
- session.headers.update(
- {
- "Authorization": f"Bearer {self.baichuan_api_key.get_secret_value()}",
- "Accept-Encoding": "identity",
- "Content-type": "application/json",
- }
- )
- self.session = session
- return self
-
- def _embed(self, texts: List[str]) -> Optional[List[List[float]]]:
- """Internal method to call Baichuan Embedding API and return embeddings.
-
- Args:
- texts: A list of texts to embed.
-
- Returns:
- A list of list of floats representing the embeddings, or None if an
- error occurs.
- """
- chunk_texts = [
- texts[i : i + self.chunk_size]
- for i in range(0, len(texts), self.chunk_size)
- ]
- embed_results = []
- for chunk in chunk_texts:
- response = self.session.post(
- BAICHUAN_API_URL, json={"input": chunk, "model": self.model_name}
- )
- # Raise exception if response status code from 400 to 600
- response.raise_for_status()
- # Check if the response status code indicates success
- if response.status_code == 200:
- resp = response.json()
- embeddings = resp.get("data", [])
- # Sort resulting embeddings by index
- sorted_embeddings = sorted(embeddings, key=lambda e: e.get("index", 0))
- # Return just the embeddings
- embed_results.extend(
- [result.get("embedding", []) for result in sorted_embeddings]
- )
- else:
- # Log error or handle unsuccessful response appropriately
- # Handle 100 <= status_code < 400, not include 200
- raise RequestException(
- f"Error: Received status code {response.status_code} from "
- "`BaichuanEmbedding` API"
- )
- return embed_results
-
- def embed_documents(self, texts: List[str]) -> Optional[List[List[float]]]: # type: ignore[override]
- """Public method to get embeddings for a list of documents.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- A list of embeddings, one for each text, or None if an error occurs.
- """
- return self._embed(texts)
-
- def embed_query(self, text: str) -> Optional[List[float]]: # type: ignore[override]
- """Public method to get embedding for a single query text.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text, or None if an error occurs.
- """
- result = self._embed([text])
- return result[0] if result is not None else None
diff --git a/libs/community/langchain_community/embeddings/baidu_qianfan_endpoint.py b/libs/community/langchain_community/embeddings/baidu_qianfan_endpoint.py
deleted file mode 100644
index aaba2f3487..0000000000
--- a/libs/community/langchain_community/embeddings/baidu_qianfan_endpoint.py
+++ /dev/null
@@ -1,186 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict, Field, SecretStr
-
-logger = logging.getLogger(__name__)
-
-
-class QianfanEmbeddingsEndpoint(BaseModel, Embeddings):
- """Baidu Qianfan Embeddings embedding models.
-
- Setup:
- To use, you should have the ``qianfan`` python package installed, and set
- environment variables ``QIANFAN_AK``, ``QIANFAN_SK``.
-
- .. code-block:: bash
-
- pip install qianfan
- export QIANFAN_AK="your-api-key"
- export QIANFAN_SK="your-secret_key"
-
- Instantiate:
- .. code-block:: python
-
- from langchain_community.embeddings import QianfanEmbeddingsEndpoint
-
- embeddings = QianfanEmbeddingsEndpoint()
-
- Embed:
- .. code-block:: python
-
- # embed the documents
- vectors = embeddings.embed_documents([text1, text2, ...])
-
- # embed the query
- vectors = embeddings.embed_query(text)
-
- # embed the documents with async
- vectors = await embeddings.aembed_documents([text1, text2, ...])
-
- # embed the query with async
- vectors = await embeddings.aembed_query(text)
- """ # noqa: E501
-
- qianfan_ak: Optional[SecretStr] = Field(default=None, alias="api_key")
- """Qianfan application apikey"""
-
- qianfan_sk: Optional[SecretStr] = Field(default=None, alias="secret_key")
- """Qianfan application secretkey"""
-
- chunk_size: int = 16
- """Chunk size when multiple texts are input"""
-
- model: Optional[str] = Field(default=None)
- """Model name
- you could get from https://cloud.baidu.com/doc/WENXINWORKSHOP/s/Nlks5zkzu
-
- for now, we support Embedding-V1 and
- - Embedding-V1 (默认模型)
- - bge-large-en
- - bge-large-zh
-
- preset models are mapping to an endpoint.
- `model` will be ignored if `endpoint` is set
- """
-
- endpoint: str = ""
- """Endpoint of the Qianfan Embedding, required if custom model used."""
-
- client: Any = None
- """Qianfan client"""
-
- init_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """init kwargs for qianfan client init, such as `query_per_second` which is
- associated with qianfan resource object to limit QPS"""
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """extra params for model invoke using with `do`."""
-
- model_config = ConfigDict(protected_namespaces=())
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """
- Validate whether qianfan_ak and qianfan_sk in the environment variables or
- configuration file are available or not.
-
- init qianfan embedding client with `ak`, `sk`, `model`, `endpoint`
-
- Args:
-
- values: a dictionary containing configuration information, must include the
- fields of qianfan_ak and qianfan_sk
- Returns:
-
- a dictionary containing configuration information. If qianfan_ak and
- qianfan_sk are not provided in the environment variables or configuration
- file,the original values will be returned; otherwise, values containing
- qianfan_ak and qianfan_sk will be returned.
- Raises:
-
- ValueError: qianfan package not found, please install it with `pip install
- qianfan`
- """
- values["qianfan_ak"] = convert_to_secret_str(
- get_from_dict_or_env(
- values,
- "qianfan_ak",
- "QIANFAN_AK",
- default="",
- )
- )
- values["qianfan_sk"] = convert_to_secret_str(
- get_from_dict_or_env(
- values,
- "qianfan_sk",
- "QIANFAN_SK",
- default="",
- )
- )
-
- try:
- import qianfan
-
- params = {
- **values.get("init_kwargs", {}),
- "model": values["model"],
- }
- if values["qianfan_ak"].get_secret_value() != "":
- params["ak"] = values["qianfan_ak"].get_secret_value()
- if values["qianfan_sk"].get_secret_value() != "":
- params["sk"] = values["qianfan_sk"].get_secret_value()
- if values["endpoint"] is not None and values["endpoint"] != "":
- params["endpoint"] = values["endpoint"]
- values["client"] = qianfan.Embedding(**params)
- except ImportError:
- raise ImportError(
- "qianfan package not found, please install it with "
- "`pip install qianfan`"
- )
- return values
-
- def embed_query(self, text: str) -> List[float]:
- resp = self.embed_documents([text])
- return resp[0]
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """
- Embeds a list of text documents using the AutoVOT algorithm.
-
- Args:
- texts (List[str]): A list of text documents to embed.
-
- Returns:
- List[List[float]]: A list of embeddings for each document in the input list.
- Each embedding is represented as a list of float values.
- """
- text_in_chunks = [
- texts[i : i + self.chunk_size]
- for i in range(0, len(texts), self.chunk_size)
- ]
- lst = []
- for chunk in text_in_chunks:
- resp = self.client.do(texts=chunk, **self.model_kwargs)
- lst.extend([res["embedding"] for res in resp["data"]])
- return lst
-
- async def aembed_query(self, text: str) -> List[float]:
- embeddings = await self.aembed_documents([text])
- return embeddings[0]
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- text_in_chunks = [
- texts[i : i + self.chunk_size]
- for i in range(0, len(texts), self.chunk_size)
- ]
- lst = []
- for chunk in text_in_chunks:
- resp = await self.client.ado(texts=chunk, **self.model_kwargs)
- for res in resp["data"]:
- lst.extend([res["embedding"]])
- return lst
diff --git a/libs/community/langchain_community/embeddings/bedrock.py b/libs/community/langchain_community/embeddings/bedrock.py
deleted file mode 100644
index 7fcfe707b2..0000000000
--- a/libs/community/langchain_community/embeddings/bedrock.py
+++ /dev/null
@@ -1,222 +0,0 @@
-import asyncio
-import json
-import os
-from typing import Any, Dict, List, Optional
-
-import numpy as np
-from langchain_core._api.deprecation import deprecated
-from langchain_core.embeddings import Embeddings
-from langchain_core.runnables.config import run_in_executor
-from pydantic import BaseModel, ConfigDict, model_validator
-from typing_extensions import Self
-
-
-@deprecated(
- since="0.2.11",
- removal="1.0",
- alternative_import="langchain_aws.BedrockEmbeddings",
-)
-class BedrockEmbeddings(BaseModel, Embeddings):
- """Bedrock embedding models.
-
- To authenticate, the AWS client uses the following methods to
- automatically load credentials:
- https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html
-
- If a specific credential profile should be used, you must pass
- the name of the profile from the ~/.aws/credentials file that is to be used.
-
- Make sure the credentials / roles used have the required policies to
- access the Bedrock service.
- """
-
- """
- Example:
- .. code-block:: python
-
- from langchain_community.bedrock_embeddings import BedrockEmbeddings
-
- region_name ="us-east-1"
- credentials_profile_name = "default"
- model_id = "amazon.titan-embed-text-v1"
-
- be = BedrockEmbeddings(
- credentials_profile_name=credentials_profile_name,
- region_name=region_name,
- model_id=model_id
- )
- """
-
- client: Any = None #: :meta private:
- """Bedrock client."""
- region_name: Optional[str] = None
- """The aws region e.g., `us-west-2`. Fallsback to AWS_DEFAULT_REGION env variable
- or region specified in ~/.aws/config in case it is not provided here.
- """
-
- credentials_profile_name: Optional[str] = None
- """The name of the profile in the ~/.aws/credentials or ~/.aws/config files, which
- has either access keys or role information specified.
- If not specified, the default credential profile or, if on an EC2 instance,
- credentials from IMDS will be used.
- See: https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html
- """
-
- model_id: str = "amazon.titan-embed-text-v1"
- """Id of the model to call, e.g., amazon.titan-embed-text-v1, this is
- equivalent to the modelId property in the list-foundation-models api"""
-
- model_kwargs: Optional[Dict] = None
- """Keyword arguments to pass to the model."""
-
- endpoint_url: Optional[str] = None
- """Needed if you don't want to default to us-east-1 endpoint"""
-
- normalize: bool = False
- """Whether the embeddings should be normalized to unit vectors"""
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- @model_validator(mode="after")
- def validate_environment(self) -> Self:
- """Validate that AWS credentials to and python package exists in environment."""
-
- if self.client is not None:
- return self
-
- try:
- import boto3
-
- if self.credentials_profile_name is not None:
- session = boto3.Session(profile_name=self.credentials_profile_name)
- else:
- # use default credentials
- session = boto3.Session()
-
- client_params = {}
- if self.region_name:
- client_params["region_name"] = self.region_name
-
- if self.endpoint_url:
- client_params["endpoint_url"] = self.endpoint_url
-
- self.client = session.client("bedrock-runtime", **client_params)
-
- except ImportError:
- raise ImportError(
- "Could not import boto3 python package. "
- "Please install it with `pip install boto3`."
- )
- except Exception as e:
- raise ValueError(
- "Could not load credentials to authenticate with AWS client. "
- "Please check that credentials in the specified "
- f"profile name are valid. Bedrock error: {e}"
- ) from e
-
- return self
-
- def _embedding_func(self, text: str) -> List[float]:
- """Call out to Bedrock embedding endpoint."""
- # replace newlines, which can negatively affect performance.
- text = text.replace(os.linesep, " ")
-
- # format input body for provider
- provider = self.model_id.split(".")[0]
- _model_kwargs = self.model_kwargs or {}
- input_body = {**_model_kwargs}
- if provider == "cohere":
- if "input_type" not in input_body.keys():
- input_body["input_type"] = "search_document"
- input_body["texts"] = [text]
- else:
- # includes common provider == "amazon"
- input_body["inputText"] = text
- body = json.dumps(input_body)
-
- try:
- # invoke bedrock API
- response = self.client.invoke_model(
- body=body,
- modelId=self.model_id,
- accept="application/json",
- contentType="application/json",
- )
-
- # format output based on provider
- response_body = json.loads(response.get("body").read())
- if provider == "cohere":
- return response_body.get("embeddings")[0]
- else:
- # includes common provider == "amazon"
- return response_body.get("embedding")
- except Exception as e:
- raise ValueError(f"Error raised by inference endpoint: {e}")
-
- def _normalize_vector(self, embeddings: List[float]) -> List[float]:
- """Normalize the embedding to a unit vector."""
- emb = np.array(embeddings)
- norm_emb = emb / np.linalg.norm(emb)
- return norm_emb.tolist()
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Compute doc embeddings using a Bedrock model.
-
- Args:
- texts: The list of texts to embed
-
- Returns:
- List of embeddings, one for each text.
- """
- results = []
- for text in texts:
- response = self._embedding_func(text)
-
- if self.normalize:
- response = self._normalize_vector(response)
-
- results.append(response)
-
- return results
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using a Bedrock model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- embedding = self._embedding_func(text)
-
- if self.normalize:
- return self._normalize_vector(embedding)
-
- return embedding
-
- async def aembed_query(self, text: str) -> List[float]:
- """Asynchronous compute query embeddings using a Bedrock model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
-
- return await run_in_executor(None, self.embed_query, text)
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- """Asynchronous compute doc embeddings using a Bedrock model.
-
- Args:
- texts: The list of texts to embed
-
- Returns:
- List of embeddings, one for each text.
- """
-
- result = await asyncio.gather(*[self.aembed_query(text) for text in texts])
-
- return list(result)
diff --git a/libs/community/langchain_community/embeddings/bookend.py b/libs/community/langchain_community/embeddings/bookend.py
deleted file mode 100644
index 76aac46fd8..0000000000
--- a/libs/community/langchain_community/embeddings/bookend.py
+++ /dev/null
@@ -1,97 +0,0 @@
-"""Wrapper around Bookend AI embedding models."""
-
-import json
-from typing import Any, List
-
-import requests
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict, Field
-
-API_URL = "https://api.bookend.ai/"
-DEFAULT_TASK = "embeddings"
-PATH = "/models/predict"
-
-
-class BookendEmbeddings(BaseModel, Embeddings):
- """Bookend AI sentence_transformers embedding models.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import BookendEmbeddings
-
- bookend = BookendEmbeddings(
- domain={domain}
- api_token={api_token}
- model_id={model_id}
- )
- bookend.embed_documents([
- "Please put on these earmuffs because I can't you hear.",
- "Baby wipes are made of chocolate stardust.",
- ])
- bookend.embed_query(
- "She only paints with bold colors; she does not like pastels."
- )
- """
-
- domain: str
- """Request for a domain at https://bookend.ai/ to use this embeddings module."""
- api_token: str
- """Request for an API token at https://bookend.ai/ to use this embeddings module."""
- model_id: str
- """Embeddings model ID to use."""
- auth_header: dict = Field(default_factory=dict)
-
- model_config = ConfigDict(protected_namespaces=())
-
- def __init__(self, **kwargs: Any):
- super().__init__(**kwargs)
- self.auth_header = {"Authorization": "Basic {}".format(self.api_token)}
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed documents using a Bookend deployed embeddings model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- result = []
- headers = self.auth_header
- headers["Content-Type"] = "application/json; charset=utf-8"
- params = {
- "model_id": self.model_id,
- "task": DEFAULT_TASK,
- }
-
- for text in texts:
- data = json.dumps(
- {
- "text": text,
- "question": None,
- "context": None,
- "instruction": None,
- }
- )
- r = requests.request(
- "POST",
- API_URL + self.domain + PATH,
- headers=headers,
- params=params,
- data=data,
- )
- result.append(r.json()[0]["data"])
-
- return result
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using a Bookend deployed embeddings model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self.embed_documents([text])[0]
diff --git a/libs/community/langchain_community/embeddings/clarifai.py b/libs/community/langchain_community/embeddings/clarifai.py
deleted file mode 100644
index e460020bef..0000000000
--- a/libs/community/langchain_community/embeddings/clarifai.py
+++ /dev/null
@@ -1,139 +0,0 @@
-import logging
-from typing import Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict, Field, model_validator
-
-logger = logging.getLogger(__name__)
-
-
-class ClarifaiEmbeddings(BaseModel, Embeddings):
- """Clarifai embedding models.
-
- To use, you should have the ``clarifai`` python package installed, and the
- environment variable ``CLARIFAI_PAT`` set with your personal access token or pass it
- as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import ClarifaiEmbeddings
- clarifai = ClarifaiEmbeddings(user_id=USER_ID,
- app_id=APP_ID,
- model_id=MODEL_ID)
- (or)
- Example_URL = "https://clarifai.com/clarifai/main/models/BAAI-bge-base-en-v15"
- clarifai = ClarifaiEmbeddings(model_url=EXAMPLE_URL)
- """
-
- model_url: Optional[str] = None
- """Model url to use."""
- model_id: Optional[str] = None
- """Model id to use."""
- model_version_id: Optional[str] = None
- """Model version id to use."""
- app_id: Optional[str] = None
- """Clarifai application id to use."""
- user_id: Optional[str] = None
- """Clarifai user id to use."""
- pat: Optional[str] = Field(default=None, exclude=True)
- """Clarifai personal access token to use."""
- token: Optional[str] = Field(default=None, exclude=True)
- """Clarifai session token to use."""
- model: Any = Field(default=None, exclude=True) #: :meta private:
- api_base: str = "https://api.clarifai.com"
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that we have all required info to access Clarifai
- platform and python package exists in environment."""
-
- try:
- from clarifai.client.model import Model
- except ImportError:
- raise ImportError(
- "Could not import clarifai python package. "
- "Please install it with `pip install clarifai`."
- )
- user_id = values.get("user_id")
- app_id = values.get("app_id")
- model_id = values.get("model_id")
- model_version_id = values.get("model_version_id")
- model_url = values.get("model_url")
- api_base = values.get("api_base")
- pat = values.get("pat")
- token = values.get("token")
-
- values["model"] = Model(
- url=model_url,
- app_id=app_id,
- user_id=user_id,
- model_version=dict(id=model_version_id),
- pat=pat,
- token=token,
- model_id=model_id,
- base_url=api_base,
- )
-
- return values
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Call out to Clarifai's embedding models.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- from clarifai.client.input import Inputs
-
- input_obj = Inputs.from_auth_helper(self.model.auth_helper)
- batch_size = 32
- embeddings = []
-
- try:
- for i in range(0, len(texts), batch_size):
- batch = texts[i : i + batch_size]
- input_batch = [
- input_obj.get_text_input(input_id=str(id), raw_text=inp)
- for id, inp in enumerate(batch)
- ]
- predict_response = self.model.predict(input_batch)
- embeddings.extend(
- [
- list(output.data.embeddings[0].vector)
- for output in predict_response.outputs
- ]
- )
-
- except Exception as e:
- logger.error(f"Predict failed, exception: {e}")
-
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to Clarifai's embedding models.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
-
- try:
- predict_response = self.model.predict_by_bytes(
- bytes(text, "utf-8"), input_type="text"
- )
- embeddings = [
- list(op.data.embeddings[0].vector) for op in predict_response.outputs
- ]
-
- except Exception as e:
- logger.error(f"Predict failed, exception: {e}")
-
- return embeddings[0]
diff --git a/libs/community/langchain_community/embeddings/cloudflare_workersai.py b/libs/community/langchain_community/embeddings/cloudflare_workersai.py
deleted file mode 100644
index 87d15c6099..0000000000
--- a/libs/community/langchain_community/embeddings/cloudflare_workersai.py
+++ /dev/null
@@ -1,91 +0,0 @@
-from typing import Any, Dict, List
-
-import requests
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict
-
-DEFAULT_MODEL_NAME = "@cf/baai/bge-base-en-v1.5"
-
-
-class CloudflareWorkersAIEmbeddings(BaseModel, Embeddings):
- """Cloudflare Workers AI embedding model.
-
- To use, you need to provide an API token and
- account ID to access Cloudflare Workers AI.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import CloudflareWorkersAIEmbeddings
-
- account_id = "my_account_id"
- api_token = "my_secret_api_token"
- model_name = "@cf/baai/bge-small-en-v1.5"
-
- cf = CloudflareWorkersAIEmbeddings(
- account_id=account_id,
- api_token=api_token,
- model_name=model_name
- )
- """
-
- api_base_url: str = "https://api.cloudflare.com/client/v4/accounts"
- account_id: str
- api_token: str
- model_name: str = DEFAULT_MODEL_NAME
- batch_size: int = 50
- strip_new_lines: bool = True
- headers: Dict[str, str] = {"Authorization": "Bearer "}
-
- def __init__(self, **kwargs: Any):
- """Initialize the Cloudflare Workers AI client."""
- super().__init__(**kwargs)
-
- self.headers = {"Authorization": f"Bearer {self.api_token}"}
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Compute doc embeddings using Cloudflare Workers AI.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- if self.strip_new_lines:
- texts = [text.replace("\n", " ") for text in texts]
-
- batches = [
- texts[i : i + self.batch_size]
- for i in range(0, len(texts), self.batch_size)
- ]
- embeddings = []
-
- for batch in batches:
- response = requests.post(
- f"{self.api_base_url}/{self.account_id}/ai/run/{self.model_name}",
- headers=self.headers,
- json={"text": batch},
- )
- embeddings.extend(response.json()["result"]["data"])
-
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using Cloudflare Workers AI.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- text = text.replace("\n", " ") if self.strip_new_lines else text
- response = requests.post(
- f"{self.api_base_url}/{self.account_id}/ai/run/{self.model_name}",
- headers=self.headers,
- json={"text": [text]},
- )
- return response.json()["result"]["data"][0]
diff --git a/libs/community/langchain_community/embeddings/clova.py b/libs/community/langchain_community/embeddings/clova.py
deleted file mode 100644
index d6d3d77b74..0000000000
--- a/libs/community/langchain_community/embeddings/clova.py
+++ /dev/null
@@ -1,142 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Dict, List, Optional, cast
-
-import requests
-from langchain_core._api.deprecation import deprecated
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env
-from pydantic import BaseModel, ConfigDict, SecretStr, model_validator
-
-
-@deprecated(
- since="0.3.4",
- removal="1.0.0",
- alternative_import="langchain_community.ClovaXEmbeddings",
-)
-class ClovaEmbeddings(BaseModel, Embeddings):
- """
- Clova's embedding service.
-
- To use this service,
-
- you should have the following environment variables
- set with your API tokens and application ID,
- or pass them as named parameters to the constructor:
-
- - ``CLOVA_EMB_API_KEY``: API key for accessing Clova's embedding service.
- - ``CLOVA_EMB_APIGW_API_KEY``: API gateway key for enhanced security.
- - ``CLOVA_EMB_APP_ID``: Application ID for identifying your application.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import ClovaEmbeddings
- embeddings = ClovaEmbeddings(
- clova_emb_api_key='your_clova_emb_api_key',
- clova_emb_apigw_api_key='your_clova_emb_apigw_api_key',
- app_id='your_app_id'
- )
-
- query_text = "This is a test query."
- query_result = embeddings.embed_query(query_text)
-
- document_text = "This is a test document."
- document_result = embeddings.embed_documents([document_text])
-
- """
-
- endpoint_url: str = (
- "https://clovastudio.apigw.ntruss.com/testapp/v1/api-tools/embedding"
- )
- """Endpoint URL to use."""
- model: str = "clir-emb-dolphin"
- """Embedding model name to use."""
- clova_emb_api_key: Optional[SecretStr] = None
- """API key for accessing Clova's embedding service."""
- clova_emb_apigw_api_key: Optional[SecretStr] = None
- """API gateway key for enhanced security."""
- app_id: Optional[SecretStr] = None
- """Application ID for identifying your application."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate api key exists in environment."""
- values["clova_emb_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(values, "clova_emb_api_key", "CLOVA_EMB_API_KEY")
- )
- values["clova_emb_apigw_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(
- values, "clova_emb_apigw_api_key", "CLOVA_EMB_APIGW_API_KEY"
- )
- )
- values["app_id"] = convert_to_secret_str(
- get_from_dict_or_env(values, "app_id", "CLOVA_EMB_APP_ID")
- )
- return values
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """
- Embed a list of texts and return their embeddings.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- embeddings = []
- for text in texts:
- embeddings.append(self._embed_text(text))
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """
- Embed a single query text and return its embedding.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self._embed_text(text)
-
- def _embed_text(self, text: str) -> List[float]:
- """
- Internal method to call the embedding API and handle the response.
- """
- payload = {"text": text}
-
- # HTTP headers for authorization
- headers = {
- "X-NCP-CLOVASTUDIO-API-KEY": cast(
- SecretStr, self.clova_emb_api_key
- ).get_secret_value(),
- "X-NCP-APIGW-API-KEY": cast(
- SecretStr, self.clova_emb_apigw_api_key
- ).get_secret_value(),
- "Content-Type": "application/json",
- }
-
- # send request
- app_id = cast(SecretStr, self.app_id).get_secret_value()
- response = requests.post(
- f"{self.endpoint_url}/{self.model}/{app_id}",
- headers=headers,
- json=payload,
- )
-
- # check for errors
- if response.status_code == 200:
- response_data = response.json()
- if "result" in response_data and "embedding" in response_data["result"]:
- return response_data["result"]["embedding"]
- raise ValueError(
- f"API request failed with status {response.status_code}: {response.text}"
- )
diff --git a/libs/community/langchain_community/embeddings/cohere.py b/libs/community/langchain_community/embeddings/cohere.py
deleted file mode 100644
index 504f688100..0000000000
--- a/libs/community/langchain_community/embeddings/cohere.py
+++ /dev/null
@@ -1,172 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import BaseModel, ConfigDict, model_validator
-
-from langchain_community.llms.cohere import _create_retry_decorator
-
-
-@deprecated(
- since="0.0.30",
- removal="1.0",
- alternative_import="langchain_cohere.CohereEmbeddings",
-)
-class CohereEmbeddings(BaseModel, Embeddings):
- """Cohere embedding models.
-
- To use, you should have the ``cohere`` python package installed, and the
- environment variable ``COHERE_API_KEY`` set with your API key or pass it
- as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import CohereEmbeddings
- cohere = CohereEmbeddings(
- model="embed-english-light-v3.0",
- cohere_api_key="my-api-key"
- )
- """
-
- client: Any = None #: :meta private:
- """Cohere client."""
- async_client: Any = None #: :meta private:
- """Cohere async client."""
- model: str = "embed-english-v2.0"
- """Model name to use."""
-
- truncate: Optional[str] = None
- """Truncate embeddings that are too long from start or end ("NONE"|"START"|"END")"""
-
- cohere_api_key: Optional[str] = None
-
- max_retries: int = 3
- """Maximum number of retries to make when generating."""
- request_timeout: Optional[float] = None
- """Timeout in seconds for the Cohere API request."""
- user_agent: str = "langchain"
- """Identifier for the application making the request."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and python package exists in environment."""
- cohere_api_key = get_from_dict_or_env(
- values, "cohere_api_key", "COHERE_API_KEY"
- )
- request_timeout = values.get("request_timeout")
-
- try:
- import cohere
-
- client_name = values["user_agent"]
- values["client"] = cohere.Client(
- cohere_api_key,
- timeout=request_timeout,
- client_name=client_name,
- )
- values["async_client"] = cohere.AsyncClient(
- cohere_api_key,
- timeout=request_timeout,
- client_name=client_name,
- )
- except ImportError:
- raise ImportError(
- "Could not import cohere python package. "
- "Please install it with `pip install cohere`."
- )
- return values
-
- def embed_with_retry(self, **kwargs: Any) -> Any:
- """Use tenacity to retry the embed call."""
- retry_decorator = _create_retry_decorator(self.max_retries)
-
- @retry_decorator
- def _embed_with_retry(**kwargs: Any) -> Any:
- return self.client.embed(**kwargs)
-
- return _embed_with_retry(**kwargs)
-
- def aembed_with_retry(self, **kwargs: Any) -> Any:
- """Use tenacity to retry the embed call."""
- retry_decorator = _create_retry_decorator(self.max_retries)
-
- @retry_decorator
- async def _embed_with_retry(**kwargs: Any) -> Any:
- return await self.async_client.embed(**kwargs)
-
- return _embed_with_retry(**kwargs)
-
- def embed(
- self, texts: List[str], *, input_type: Optional[str] = None
- ) -> List[List[float]]:
- embeddings = self.embed_with_retry(
- model=self.model,
- texts=texts,
- input_type=input_type,
- truncate=self.truncate,
- ).embeddings
- return [list(map(float, e)) for e in embeddings]
-
- async def aembed(
- self, texts: List[str], *, input_type: Optional[str] = None
- ) -> List[List[float]]:
- embeddings = (
- await self.aembed_with_retry(
- model=self.model,
- texts=texts,
- input_type=input_type,
- truncate=self.truncate,
- )
- ).embeddings
- return [list(map(float, e)) for e in embeddings]
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed a list of document texts.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- return self.embed(texts, input_type="search_document")
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- """Async call out to Cohere's embedding endpoint.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- return await self.aembed(texts, input_type="search_document")
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to Cohere's embedding endpoint.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self.embed([text], input_type="search_query")[0]
-
- async def aembed_query(self, text: str) -> List[float]:
- """Async call out to Cohere's embedding endpoint.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return (await self.aembed([text], input_type="search_query"))[0]
diff --git a/libs/community/langchain_community/embeddings/dashscope.py b/libs/community/langchain_community/embeddings/dashscope.py
deleted file mode 100644
index 59bd76e0de..0000000000
--- a/libs/community/langchain_community/embeddings/dashscope.py
+++ /dev/null
@@ -1,168 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import (
- Any,
- Callable,
- Dict,
- List,
- Optional,
-)
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import BaseModel, ConfigDict, model_validator
-from requests.exceptions import HTTPError
-from tenacity import (
- before_sleep_log,
- retry,
- retry_if_exception_type,
- stop_after_attempt,
- wait_exponential,
-)
-
-logger = logging.getLogger(__name__)
-
-BATCH_SIZE = {"text-embedding-v1": 25, "text-embedding-v2": 25, "text-embedding-v3": 6}
-
-
-def _create_retry_decorator(embeddings: DashScopeEmbeddings) -> Callable[[Any], Any]:
- multiplier = 1
- min_seconds = 1
- max_seconds = 4
- # Wait 2^x * 1 second between each retry starting with
- # 1 seconds, then up to 4 seconds, then 4 seconds afterwards
- return retry(
- reraise=True,
- stop=stop_after_attempt(embeddings.max_retries),
- wait=wait_exponential(multiplier, min=min_seconds, max=max_seconds),
- retry=(retry_if_exception_type(HTTPError)),
- before_sleep=before_sleep_log(logger, logging.WARNING),
- )
-
-
-def embed_with_retry(embeddings: DashScopeEmbeddings, **kwargs: Any) -> Any:
- """Use tenacity to retry the embedding call."""
- retry_decorator = _create_retry_decorator(embeddings)
-
- @retry_decorator
- def _embed_with_retry(**kwargs: Any) -> Any:
- result = []
- i = 0
- input_data = kwargs["input"]
- input_len = len(input_data) if isinstance(input_data, list) else 1
- batch_size = BATCH_SIZE.get(kwargs["model"], 25)
- while i < input_len:
- kwargs["input"] = (
- input_data[i : i + batch_size]
- if isinstance(input_data, list)
- else input_data
- )
- resp = embeddings.client.call(**kwargs)
- if resp.status_code == 200:
- result += resp.output["embeddings"]
- elif resp.status_code in [400, 401]:
- raise ValueError(
- f"status_code: {resp.status_code} \n "
- f"code: {resp.code} \n message: {resp.message}"
- )
- else:
- raise HTTPError(
- f"HTTP error occurred: status_code: {resp.status_code} \n "
- f"code: {resp.code} \n message: {resp.message}",
- response=resp,
- )
- i += batch_size
- return result
-
- return _embed_with_retry(**kwargs)
-
-
-class DashScopeEmbeddings(BaseModel, Embeddings):
- """DashScope embedding models.
-
- To use, you should have the ``dashscope`` python package installed, and the
- environment variable ``DASHSCOPE_API_KEY`` set with your API key or pass it
- as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import DashScopeEmbeddings
- embeddings = DashScopeEmbeddings(dashscope_api_key="my-api-key")
-
- Example:
- .. code-block:: python
-
- import os
- os.environ["DASHSCOPE_API_KEY"] = "your DashScope API KEY"
-
- from langchain_community.embeddings.dashscope import DashScopeEmbeddings
- embeddings = DashScopeEmbeddings(
- model="text-embedding-v1",
- )
- text = "This is a test query."
- query_result = embeddings.embed_query(text)
-
- """
-
- client: Any = None #: :meta private:
- """The DashScope client."""
- model: str = "text-embedding-v1"
- dashscope_api_key: Optional[str] = None
- max_retries: int = 5
- """Maximum number of retries to make when generating."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- import dashscope
-
- """Validate that api key and python package exists in environment."""
- values["dashscope_api_key"] = get_from_dict_or_env(
- values, "dashscope_api_key", "DASHSCOPE_API_KEY"
- )
- dashscope.api_key = values["dashscope_api_key"]
- try:
- import dashscope
-
- values["client"] = dashscope.TextEmbedding
- except ImportError:
- raise ImportError(
- "Could not import dashscope python package. "
- "Please install it with `pip install dashscope`."
- )
- return values
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Call out to DashScope's embedding endpoint for embedding search docs.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- embeddings = embed_with_retry(
- self, input=texts, text_type="document", model=self.model
- )
- embedding_list = [item["embedding"] for item in embeddings]
- return embedding_list
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to DashScope's embedding endpoint for embedding query text.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embedding for the text.
- """
- embedding = embed_with_retry(
- self, input=text, text_type="query", model=self.model
- )[0]["embedding"]
- return embedding
diff --git a/libs/community/langchain_community/embeddings/databricks.py b/libs/community/langchain_community/embeddings/databricks.py
deleted file mode 100644
index 2bb68024b5..0000000000
--- a/libs/community/langchain_community/embeddings/databricks.py
+++ /dev/null
@@ -1,52 +0,0 @@
-from __future__ import annotations
-
-from typing import Iterator, List
-from urllib.parse import urlparse
-
-from langchain_core._api import deprecated
-
-from langchain_community.embeddings.mlflow import MlflowEmbeddings
-
-
-def _chunk(texts: List[str], size: int) -> Iterator[List[str]]:
- for i in range(0, len(texts), size):
- yield texts[i : i + size]
-
-
-@deprecated(
- since="0.3.3",
- removal="1.0",
- alternative_import="databricks_langchain.DatabricksEmbeddings",
-)
-class DatabricksEmbeddings(MlflowEmbeddings):
- """Databricks embeddings.
-
- To use, you should have the ``mlflow`` python package installed.
- For more information, see https://mlflow.org/docs/latest/llms/deployments.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import DatabricksEmbeddings
-
- embeddings = DatabricksEmbeddings(
- target_uri="databricks",
- endpoint="embeddings",
- )
- """
-
- target_uri: str = "databricks"
- """The target URI to use. Defaults to ``databricks``."""
-
- @property
- def _mlflow_extras(self) -> str:
- return ""
-
- def _validate_uri(self) -> None:
- if self.target_uri == "databricks":
- return
-
- if urlparse(self.target_uri).scheme != "databricks":
- raise ValueError(
- "Invalid target URI. The target URI must be a valid databricks URI."
- )
diff --git a/libs/community/langchain_community/embeddings/deepinfra.py b/libs/community/langchain_community/embeddings/deepinfra.py
deleted file mode 100644
index d0d2c47601..0000000000
--- a/libs/community/langchain_community/embeddings/deepinfra.py
+++ /dev/null
@@ -1,140 +0,0 @@
-from typing import Any, Dict, List, Mapping, Optional
-
-import requests
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict
-
-DEFAULT_MODEL_ID = "sentence-transformers/clip-ViT-B-32"
-MAX_BATCH_SIZE = 1024
-
-
-class DeepInfraEmbeddings(BaseModel, Embeddings):
- """Deep Infra's embedding inference service.
-
- To use, you should have the
- environment variable ``DEEPINFRA_API_TOKEN`` set with your API token, or pass
- it as a named parameter to the constructor.
- There are multiple embeddings models available,
- see https://deepinfra.com/models?type=embeddings.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import DeepInfraEmbeddings
- deepinfra_emb = DeepInfraEmbeddings(
- model_id="sentence-transformers/clip-ViT-B-32",
- deepinfra_api_token="my-api-key"
- )
- r1 = deepinfra_emb.embed_documents(
- [
- "Alpha is the first letter of Greek alphabet",
- "Beta is the second letter of Greek alphabet",
- ]
- )
- r2 = deepinfra_emb.embed_query(
- "What is the second letter of Greek alphabet"
- )
-
- """
-
- model_id: str = DEFAULT_MODEL_ID
- """Embeddings model to use."""
- normalize: bool = False
- """whether to normalize the computed embeddings"""
- embed_instruction: str = "passage: "
- """Instruction used to embed documents."""
- query_instruction: str = "query: "
- """Instruction used to embed the query."""
- model_kwargs: Optional[dict] = None
- """Other model keyword args"""
- deepinfra_api_token: Optional[str] = None
- """API token for Deep Infra. If not provided, the token is
- fetched from the environment variable 'DEEPINFRA_API_TOKEN'."""
- batch_size: int = MAX_BATCH_SIZE
- """Batch size for embedding requests."""
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- deepinfra_api_token = get_from_dict_or_env(
- values, "deepinfra_api_token", "DEEPINFRA_API_TOKEN"
- )
- values["deepinfra_api_token"] = deepinfra_api_token
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {"model_id": self.model_id}
-
- def _embed(self, input: List[str]) -> List[List[float]]:
- _model_kwargs = self.model_kwargs or {}
- # HTTP headers for authorization
- headers = {
- "Authorization": f"bearer {self.deepinfra_api_token}",
- "Content-Type": "application/json",
- }
- # send request
- try:
- res = requests.post(
- f"https://api.deepinfra.com/v1/inference/{self.model_id}",
- headers=headers,
- json={"inputs": input, "normalize": self.normalize, **_model_kwargs},
- )
- except requests.exceptions.RequestException as e:
- raise ValueError(f"Error raised by inference endpoint: {e}")
-
- if res.status_code != 200:
- raise ValueError(
- "Error raised by inference API HTTP code: %s, %s"
- % (res.status_code, res.text)
- )
- try:
- t = res.json()
- embeddings = t["embeddings"]
- except requests.exceptions.JSONDecodeError as e:
- raise ValueError(
- f"Error raised by inference API: {e}.\nResponse: {res.text}"
- )
-
- return embeddings
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed documents using a Deep Infra deployed embedding model.
- For larger batches, the input list of texts is chunked into smaller
- batches to avoid exceeding the maximum request size.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
-
- embeddings = []
- instruction_pairs = [f"{self.embed_instruction}{text}" for text in texts]
-
- chunks = [
- instruction_pairs[i : i + self.batch_size]
- for i in range(0, len(instruction_pairs), self.batch_size)
- ]
- for chunk in chunks:
- embeddings += self._embed(chunk)
-
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using a Deep Infra deployed embedding model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- instruction_pair = f"{self.query_instruction}{text}"
- embedding = self._embed([instruction_pair])[0]
- return embedding
diff --git a/libs/community/langchain_community/embeddings/edenai.py b/libs/community/langchain_community/embeddings/edenai.py
deleted file mode 100644
index 097c730ae4..0000000000
--- a/libs/community/langchain_community/embeddings/edenai.py
+++ /dev/null
@@ -1,114 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import (
- BaseModel,
- ConfigDict,
- Field,
- SecretStr,
-)
-
-from langchain_community.utilities.requests import Requests
-
-
-class EdenAiEmbeddings(BaseModel, Embeddings):
- """EdenAI embedding.
- environment variable ``EDENAI_API_KEY`` set with your API key, or pass
- it as a named parameter.
- """
-
- edenai_api_key: Optional[SecretStr] = Field(None, description="EdenAI API Token")
-
- provider: str = "openai"
- """embedding provider to use (eg: openai,google etc.)"""
-
- model: Optional[str] = None
- """
- model name for above provider (eg: 'gpt-3.5-turbo-instruct' for openai)
- available models are shown on https://docs.edenai.co/ under 'available providers'
- """
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key exists in environment."""
- values["edenai_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(values, "edenai_api_key", "EDENAI_API_KEY")
- )
- return values
-
- @staticmethod
- def get_user_agent() -> str:
- from langchain_community import __version__
-
- return f"langchain/{__version__}"
-
- def _generate_embeddings(self, texts: List[str]) -> List[List[float]]:
- """Compute embeddings using EdenAi api."""
- url = "https://api.edenai.run/v2/text/embeddings"
-
- headers = {
- "accept": "application/json",
- "content-type": "application/json",
- "authorization": f"Bearer {self.edenai_api_key.get_secret_value()}", # type: ignore[union-attr]
- "User-Agent": self.get_user_agent(),
- }
-
- payload: Dict[str, Any] = {"texts": texts, "providers": self.provider}
-
- if self.model is not None:
- payload["settings"] = {self.provider: self.model}
-
- request = Requests(headers=headers)
- response = request.post(url=url, data=payload)
- if response.status_code >= 500:
- raise Exception(f"EdenAI Server: Error {response.status_code}")
- elif response.status_code >= 400:
- raise ValueError(f"EdenAI received an invalid payload: {response.text}")
- elif response.status_code != 200:
- raise Exception(
- f"EdenAI returned an unexpected response with status "
- f"{response.status_code}: {response.text}"
- )
-
- temp = response.json()
-
- provider_response = temp[self.provider]
- if provider_response.get("status") == "fail":
- err_msg = provider_response.get("error", {}).get("message")
- raise Exception(err_msg)
-
- embeddings = []
- for embed_item in temp[self.provider]["items"]:
- embedding = embed_item["embedding"]
-
- embeddings.append(embedding)
-
- return embeddings
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed a list of documents using EdenAI.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
-
- return self._generate_embeddings(texts)
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using EdenAI.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self._generate_embeddings([text])[0]
diff --git a/libs/community/langchain_community/embeddings/elasticsearch.py b/libs/community/langchain_community/embeddings/elasticsearch.py
deleted file mode 100644
index ea080ab9aa..0000000000
--- a/libs/community/langchain_community/embeddings/elasticsearch.py
+++ /dev/null
@@ -1,226 +0,0 @@
-from __future__ import annotations
-
-from typing import TYPE_CHECKING, List, Optional
-
-from langchain_core._api import deprecated
-from langchain_core.utils import get_from_env
-
-if TYPE_CHECKING:
- from elasticsearch import Elasticsearch
- from elasticsearch.client import MlClient
-
-from langchain_core.embeddings import Embeddings
-
-
-@deprecated(
- "0.1.11", alternative="Use class in langchain-elasticsearch package", pending=True
-)
-class ElasticsearchEmbeddings(Embeddings):
- """Elasticsearch embedding models.
-
- This class provides an interface to generate embeddings using a model deployed
- in an Elasticsearch cluster. It requires an Elasticsearch connection object
- and the model_id of the model deployed in the cluster.
-
- In Elasticsearch you need to have an embedding model loaded and deployed.
- - https://www.elastic.co/guide/en/elasticsearch/reference/current/infer-trained-model.html
- - https://www.elastic.co/guide/en/machine-learning/current/ml-nlp-deploy-models.html
- """
-
- def __init__(
- self,
- client: MlClient,
- model_id: str,
- *,
- input_field: str = "text_field",
- ):
- """
- Initialize the ElasticsearchEmbeddings instance.
-
- Args:
- client (MlClient): An Elasticsearch ML client object.
- model_id (str): The model_id of the model deployed in the Elasticsearch
- cluster.
- input_field (str): The name of the key for the input text field in the
- document. Defaults to 'text_field'.
- """
- self.client = client
- self.model_id = model_id
- self.input_field = input_field
-
- @classmethod
- def from_credentials(
- cls,
- model_id: str,
- *,
- es_cloud_id: Optional[str] = None,
- es_user: Optional[str] = None,
- es_password: Optional[str] = None,
- input_field: str = "text_field",
- ) -> ElasticsearchEmbeddings:
- """Instantiate embeddings from Elasticsearch credentials.
-
- Args:
- model_id (str): The model_id of the model deployed in the Elasticsearch
- cluster.
- input_field (str): The name of the key for the input text field in the
- document. Defaults to 'text_field'.
- es_cloud_id: (str, optional): The Elasticsearch cloud ID to connect to.
- es_user: (str, optional): Elasticsearch username.
- es_password: (str, optional): Elasticsearch password.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import ElasticsearchEmbeddings
-
- # Define the model ID and input field name (if different from default)
- model_id = "your_model_id"
- # Optional, only if different from 'text_field'
- input_field = "your_input_field"
-
- # Credentials can be passed in two ways. Either set the env vars
- # ES_CLOUD_ID, ES_USER, ES_PASSWORD and they will be automatically
- # pulled in, or pass them in directly as kwargs.
- embeddings = ElasticsearchEmbeddings.from_credentials(
- model_id,
- input_field=input_field,
- # es_cloud_id="foo",
- # es_user="bar",
- # es_password="baz",
- )
-
- documents = [
- "This is an example document.",
- "Another example document to generate embeddings for.",
- ]
- embeddings_generator.embed_documents(documents)
- """
- try:
- from elasticsearch import Elasticsearch
- from elasticsearch.client import MlClient
- except ImportError:
- raise ImportError(
- "elasticsearch package not found, please install with 'pip install "
- "elasticsearch'"
- )
-
- es_cloud_id = es_cloud_id or get_from_env("es_cloud_id", "ES_CLOUD_ID")
- es_user = es_user or get_from_env("es_user", "ES_USER")
- es_password = es_password or get_from_env("es_password", "ES_PASSWORD")
-
- # Connect to Elasticsearch
- es_connection = Elasticsearch(
- cloud_id=es_cloud_id, basic_auth=(es_user, es_password)
- )
- client = MlClient(es_connection)
- return cls(client, model_id, input_field=input_field)
-
- @classmethod
- def from_es_connection(
- cls,
- model_id: str,
- es_connection: Elasticsearch,
- input_field: str = "text_field",
- ) -> ElasticsearchEmbeddings:
- """
- Instantiate embeddings from an existing Elasticsearch connection.
-
- This method provides a way to create an instance of the ElasticsearchEmbeddings
- class using an existing Elasticsearch connection. The connection object is used
- to create an MlClient, which is then used to initialize the
- ElasticsearchEmbeddings instance.
-
- Args:
- model_id (str): The model_id of the model deployed in the Elasticsearch cluster.
- es_connection (elasticsearch.Elasticsearch): An existing Elasticsearch
- connection object. input_field (str, optional): The name of the key for the
- input text field in the document. Defaults to 'text_field'.
-
- Returns:
- ElasticsearchEmbeddings: An instance of the ElasticsearchEmbeddings class.
-
- Example:
- .. code-block:: python
-
- from elasticsearch import Elasticsearch
-
- from langchain_community.embeddings import ElasticsearchEmbeddings
-
- # Define the model ID and input field name (if different from default)
- model_id = "your_model_id"
- # Optional, only if different from 'text_field'
- input_field = "your_input_field"
-
- # Create Elasticsearch connection
- es_connection = Elasticsearch(
- hosts=["localhost:9200"], http_auth=("user", "password")
- )
-
- # Instantiate ElasticsearchEmbeddings using the existing connection
- embeddings = ElasticsearchEmbeddings.from_es_connection(
- model_id,
- es_connection,
- input_field=input_field,
- )
-
- documents = [
- "This is an example document.",
- "Another example document to generate embeddings for.",
- ]
- embeddings_generator.embed_documents(documents)
- """
- # Importing MlClient from elasticsearch.client within the method to
- # avoid unnecessary import if the method is not used
- from elasticsearch.client import MlClient
-
- # Create an MlClient from the given Elasticsearch connection
- client = MlClient(es_connection)
-
- # Return a new instance of the ElasticsearchEmbeddings class with
- # the MlClient, model_id, and input_field
- return cls(client, model_id, input_field=input_field)
-
- def _embedding_func(self, texts: List[str]) -> List[List[float]]:
- """
- Generate embeddings for the given texts using the Elasticsearch model.
-
- Args:
- texts (List[str]): A list of text strings to generate embeddings for.
-
- Returns:
- List[List[float]]: A list of embeddings, one for each text in the input
- list.
- """
- response = self.client.infer_trained_model(
- model_id=self.model_id, docs=[{self.input_field: text} for text in texts]
- )
-
- embeddings = [doc["predicted_value"] for doc in response["inference_results"]]
- return embeddings
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """
- Generate embeddings for a list of documents.
-
- Args:
- texts (List[str]): A list of document text strings to generate embeddings
- for.
-
- Returns:
- List[List[float]]: A list of embeddings, one for each document in the input
- list.
- """
- return self._embedding_func(texts)
-
- def embed_query(self, text: str) -> List[float]:
- """
- Generate an embedding for a single query text.
-
- Args:
- text (str): The query text to generate an embedding for.
-
- Returns:
- List[float]: The embedding for the input query text.
- """
- return self._embedding_func([text])[0]
diff --git a/libs/community/langchain_community/embeddings/embaas.py b/libs/community/langchain_community/embeddings/embaas.py
deleted file mode 100644
index 78fd42bf85..0000000000
--- a/libs/community/langchain_community/embeddings/embaas.py
+++ /dev/null
@@ -1,155 +0,0 @@
-from typing import Any, Dict, List, Mapping, Optional
-
-import requests
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict, SecretStr
-from requests.adapters import HTTPAdapter, Retry
-from typing_extensions import NotRequired, TypedDict
-
-# Currently supported maximum batch size for embedding requests
-MAX_BATCH_SIZE = 256
-EMBAAS_API_URL = "https://api.embaas.io/v1/embeddings/"
-
-
-class EmbaasEmbeddingsPayload(TypedDict):
- """Payload for the Embaas embeddings API."""
-
- model: str
- texts: List[str]
- instruction: NotRequired[str]
-
-
-class EmbaasEmbeddings(BaseModel, Embeddings):
- """Embaas's embedding service.
-
- To use, you should have the
- environment variable ``EMBAAS_API_KEY`` set with your API key, or pass
- it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- # initialize with default model and instruction
- from langchain_community.embeddings import EmbaasEmbeddings
- emb = EmbaasEmbeddings()
-
- # initialize with custom model and instruction
- from langchain_community.embeddings import EmbaasEmbeddings
- emb_model = "instructor-large"
- emb_inst = "Represent the Wikipedia document for retrieval"
- emb = EmbaasEmbeddings(
- model=emb_model,
- instruction=emb_inst
- )
- """
-
- model: str = "e5-large-v2"
- """The model used for embeddings."""
- instruction: Optional[str] = None
- """Instruction used for domain-specific embeddings."""
- api_url: str = EMBAAS_API_URL
- """The URL for the embaas embeddings API."""
- embaas_api_key: Optional[SecretStr] = None
- """max number of retries for requests"""
- max_retries: Optional[int] = 3
- """request timeout in seconds"""
- timeout: Optional[int] = 30
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- embaas_api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "embaas_api_key", "EMBAAS_API_KEY")
- )
- values["embaas_api_key"] = embaas_api_key
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying params."""
- return {"model": self.model, "instruction": self.instruction}
-
- def _generate_payload(self, texts: List[str]) -> EmbaasEmbeddingsPayload:
- """Generates payload for the API request."""
- payload = EmbaasEmbeddingsPayload(texts=texts, model=self.model)
- if self.instruction:
- payload["instruction"] = self.instruction
- return payload
-
- def _handle_request(self, payload: EmbaasEmbeddingsPayload) -> List[List[float]]:
- """Sends a request to the Embaas API and handles the response."""
- headers = {
- "Authorization": f"Bearer {self.embaas_api_key.get_secret_value()}", # type: ignore[union-attr]
- "Content-Type": "application/json",
- }
-
- session = requests.Session()
- retries = Retry(
- total=self.max_retries,
- backoff_factor=0.5,
- allowed_methods=["POST"],
- raise_on_status=True,
- )
-
- session.mount("http://", HTTPAdapter(max_retries=retries))
- session.mount("https://", HTTPAdapter(max_retries=retries))
- response = session.post(
- self.api_url,
- headers=headers,
- json=payload,
- timeout=self.timeout,
- )
-
- parsed_response = response.json()
- embeddings = [item["embedding"] for item in parsed_response["data"]]
-
- return embeddings
-
- def _generate_embeddings(self, texts: List[str]) -> List[List[float]]:
- """Generate embeddings using the Embaas API."""
- payload = self._generate_payload(texts)
- try:
- return self._handle_request(payload)
- except requests.exceptions.RequestException as e:
- if e.response is None or not e.response.text:
- raise ValueError(f"Error raised by embaas embeddings API: {e}")
-
- parsed_response = e.response.json()
- if "message" in parsed_response:
- raise ValueError(
- "Validation Error raised by embaas embeddings API:"
- f"{parsed_response['message']}"
- )
- raise
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Get embeddings for a list of texts.
-
- Args:
- texts: The list of texts to get embeddings for.
-
- Returns:
- List of embeddings, one for each text.
- """
- batches = [
- texts[i : i + MAX_BATCH_SIZE] for i in range(0, len(texts), MAX_BATCH_SIZE)
- ]
- embeddings = [self._generate_embeddings(batch) for batch in batches]
- # flatten the list of lists into a single list
- return [embedding for batch in embeddings for embedding in batch]
-
- def embed_query(self, text: str) -> List[float]:
- """Get embeddings for a single text.
-
- Args:
- text: The text to get embeddings for.
-
- Returns:
- List of embeddings.
- """
- return self.embed_documents([text])[0]
diff --git a/libs/community/langchain_community/embeddings/ernie.py b/libs/community/langchain_community/embeddings/ernie.py
deleted file mode 100644
index 34758c58b4..0000000000
--- a/libs/community/langchain_community/embeddings/ernie.py
+++ /dev/null
@@ -1,158 +0,0 @@
-import asyncio
-import logging
-import threading
-from typing import Dict, List, Optional
-
-import requests
-from langchain_core._api.deprecation import deprecated
-from langchain_core.embeddings import Embeddings
-from langchain_core.runnables.config import run_in_executor
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict
-
-logger = logging.getLogger(__name__)
-
-
-@deprecated(
- since="0.0.13",
- alternative="langchain_community.embeddings.QianfanEmbeddingsEndpoint",
-)
-class ErnieEmbeddings(BaseModel, Embeddings):
- """`Ernie Embeddings V1` embedding models."""
-
- ernie_api_base: Optional[str] = None
- ernie_client_id: Optional[str] = None
- ernie_client_secret: Optional[str] = None
- access_token: Optional[str] = None
-
- chunk_size: int = 16
-
- model_name: str = "ErnieBot-Embedding-V1"
-
- _lock = threading.Lock()
-
- model_config = ConfigDict(protected_namespaces=())
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- values["ernie_api_base"] = get_from_dict_or_env(
- values, "ernie_api_base", "ERNIE_API_BASE", "https://aip.baidubce.com"
- )
- values["ernie_client_id"] = get_from_dict_or_env(
- values,
- "ernie_client_id",
- "ERNIE_CLIENT_ID",
- )
- values["ernie_client_secret"] = get_from_dict_or_env(
- values,
- "ernie_client_secret",
- "ERNIE_CLIENT_SECRET",
- )
- return values
-
- def _embedding(self, json: object) -> dict:
- base_url = (
- f"{self.ernie_api_base}/rpc/2.0/ai_custom/v1/wenxinworkshop/embeddings"
- )
- resp = requests.post(
- f"{base_url}/embedding-v1",
- headers={
- "Content-Type": "application/json",
- },
- params={"access_token": self.access_token},
- json=json,
- )
- return resp.json()
-
- def _refresh_access_token_with_lock(self) -> None:
- with self._lock:
- logger.debug("Refreshing access token")
- base_url: str = f"{self.ernie_api_base}/oauth/2.0/token"
- resp = requests.post(
- base_url,
- headers={
- "Content-Type": "application/json",
- "Accept": "application/json",
- },
- params={
- "grant_type": "client_credentials",
- "client_id": self.ernie_client_id,
- "client_secret": self.ernie_client_secret,
- },
- )
- self.access_token = str(resp.json().get("access_token"))
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed search docs.
-
- Args:
- texts: The list of texts to embed
-
- Returns:
- List[List[float]]: List of embeddings, one for each text.
- """
-
- if not self.access_token:
- self._refresh_access_token_with_lock()
- text_in_chunks = [
- texts[i : i + self.chunk_size]
- for i in range(0, len(texts), self.chunk_size)
- ]
- lst = []
- for chunk in text_in_chunks:
- resp = self._embedding({"input": [text for text in chunk]})
- if resp.get("error_code"):
- if resp.get("error_code") == 111:
- self._refresh_access_token_with_lock()
- resp = self._embedding({"input": [text for text in chunk]})
- else:
- raise ValueError(f"Error from Ernie: {resp}")
- lst.extend([i["embedding"] for i in resp["data"]])
- return lst
-
- def embed_query(self, text: str) -> List[float]:
- """Embed query text.
-
- Args:
- text: The text to embed.
-
- Returns:
- List[float]: Embeddings for the text.
- """
-
- if not self.access_token:
- self._refresh_access_token_with_lock()
- resp = self._embedding({"input": [text]})
- if resp.get("error_code"):
- if resp.get("error_code") == 111:
- self._refresh_access_token_with_lock()
- resp = self._embedding({"input": [text]})
- else:
- raise ValueError(f"Error from Ernie: {resp}")
- return resp["data"][0]["embedding"]
-
- async def aembed_query(self, text: str) -> List[float]:
- """Asynchronous Embed query text.
-
- Args:
- text: The text to embed.
-
- Returns:
- List[float]: Embeddings for the text.
- """
-
- return await run_in_executor(None, self.embed_query, text)
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- """Asynchronous Embed search docs.
-
- Args:
- texts: The list of texts to embed
-
- Returns:
- List[List[float]]: List of embeddings, one for each text.
- """
-
- result = await asyncio.gather(*[self.aembed_query(text) for text in texts])
-
- return list(result)
diff --git a/libs/community/langchain_community/embeddings/fake.py b/libs/community/langchain_community/embeddings/fake.py
deleted file mode 100644
index 6bbfeeb45c..0000000000
--- a/libs/community/langchain_community/embeddings/fake.py
+++ /dev/null
@@ -1,50 +0,0 @@
-import hashlib
-from typing import List
-
-import numpy as np
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel
-
-
-class FakeEmbeddings(Embeddings, BaseModel):
- """Fake embedding model."""
-
- size: int
- """The size of the embedding vector."""
-
- def _get_embedding(self) -> List[float]:
- return list(np.random.normal(size=self.size))
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- return [self._get_embedding() for _ in texts]
-
- def embed_query(self, text: str) -> List[float]:
- return self._get_embedding()
-
-
-class DeterministicFakeEmbedding(Embeddings, BaseModel):
- """
- Fake embedding model that always returns
- the same embedding vector for the same text.
- """
-
- size: int
- """The size of the embedding vector."""
-
- def _get_embedding(self, seed: int) -> List[float]:
- # set the seed for the random generator
- np.random.seed(seed)
- return list(np.random.normal(size=self.size))
-
- @staticmethod
- def _get_seed(text: str) -> int:
- """
- Get a seed for the random generator, using the hash of the text.
- """
- return int(hashlib.sha256(text.encode("utf-8")).hexdigest(), 16) % 10**8
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- return [self._get_embedding(seed=self._get_seed(_)) for _ in texts]
-
- def embed_query(self, text: str) -> List[float]:
- return self._get_embedding(seed=self._get_seed(text))
diff --git a/libs/community/langchain_community/embeddings/fastembed.py b/libs/community/langchain_community/embeddings/fastembed.py
deleted file mode 100644
index d46f921060..0000000000
--- a/libs/community/langchain_community/embeddings/fastembed.py
+++ /dev/null
@@ -1,152 +0,0 @@
-import importlib
-import importlib.metadata
-from typing import Any, Dict, List, Literal, Optional, Sequence, cast
-
-import numpy as np
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import pre_init
-from pydantic import BaseModel, ConfigDict
-
-MIN_VERSION = "0.2.0"
-
-
-class FastEmbedEmbeddings(BaseModel, Embeddings):
- """Qdrant FastEmbedding models.
-
- FastEmbed is a lightweight, fast, Python library built for embedding generation.
- See more documentation at:
- * https://github.com/qdrant/fastembed/
- * https://qdrant.github.io/fastembed/
-
- To use this class, you must install the `fastembed` Python package.
-
- `pip install fastembed`
- Example:
- from langchain_community.embeddings import FastEmbedEmbeddings
- fastembed = FastEmbedEmbeddings()
- """
-
- model_name: str = "BAAI/bge-small-en-v1.5"
- """Name of the FastEmbedding model to use
- Defaults to "BAAI/bge-small-en-v1.5"
- Find the list of supported models at
- https://qdrant.github.io/fastembed/examples/Supported_Models/
- """
-
- max_length: int = 512
- """The maximum number of tokens. Defaults to 512.
- Unknown behavior for values > 512.
- """
-
- cache_dir: Optional[str] = None
- """The path to the cache directory.
- Defaults to `local_cache` in the parent directory
- """
-
- threads: Optional[int] = None
- """The number of threads single onnxruntime session can use.
- Defaults to None
- """
-
- doc_embed_type: Literal["default", "passage"] = "default"
- """Type of embedding to use for documents
- The available options are: "default" and "passage"
- """
-
- batch_size: int = 256
- """Batch size for encoding. Higher values will use more memory, but be faster.
- Defaults to 256.
- """
-
- parallel: Optional[int] = None
- """If `>1`, parallel encoding is used, recommended for encoding of large datasets.
- If `0`, use all available cores.
- If `None`, don't use data-parallel processing, use default onnxruntime threading.
- Defaults to `None`.
- """
-
- providers: Optional[Sequence[Any]] = None
- """List of ONNX execution providers. Use `["CUDAExecutionProvider"]` to enable the
- use of GPU when generating embeddings. This requires to install `fastembed-gpu`
- instead of `fastembed`. See https://qdrant.github.io/fastembed/examples/FastEmbed_GPU
- for more details.
- Defaults to `None`.
- """
-
- model: Any = None # : :meta private:
-
- model_config = ConfigDict(extra="allow", protected_namespaces=())
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that FastEmbed has been installed."""
- model_name = values.get("model_name")
- max_length = values.get("max_length")
- cache_dir = values.get("cache_dir")
- threads = values.get("threads")
- providers = values.get("providers")
- pkg_to_install = (
- "fastembed-gpu"
- if providers and "CUDAExecutionProvider" in providers
- else "fastembed"
- )
-
- try:
- fastembed = importlib.import_module("fastembed")
-
- except ModuleNotFoundError:
- raise ImportError(
- "Could not import 'fastembed' Python package. "
- f"Please install it with `pip install {pkg_to_install}`."
- )
-
- if importlib.metadata.version(pkg_to_install) < MIN_VERSION:
- raise ImportError(
- f"FastEmbedEmbeddings requires "
- f'`pip install -U "{pkg_to_install}>={MIN_VERSION}"`.'
- )
-
- values["model"] = fastembed.TextEmbedding(
- model_name=model_name,
- max_length=max_length,
- cache_dir=cache_dir,
- threads=threads,
- providers=providers,
- )
- return values
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Generate embeddings for documents using FastEmbed.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- embeddings: List[np.ndarray]
- if self.doc_embed_type == "passage":
- embeddings = self.model.passage_embed(
- texts, batch_size=self.batch_size, parallel=self.parallel
- )
- else:
- embeddings = self.model.embed(
- texts, batch_size=self.batch_size, parallel=self.parallel
- )
- return [cast(List[float], e.tolist()) for e in embeddings]
-
- def embed_query(self, text: str) -> List[float]:
- """Generate query embeddings using FastEmbed.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- query_embeddings: np.ndarray = next(
- self.model.query_embed(
- text, batch_size=self.batch_size, parallel=self.parallel
- )
- )
- return cast(List[float], query_embeddings.tolist())
diff --git a/libs/community/langchain_community/embeddings/gigachat.py b/libs/community/langchain_community/embeddings/gigachat.py
deleted file mode 100644
index 359addfd4f..0000000000
--- a/libs/community/langchain_community/embeddings/gigachat.py
+++ /dev/null
@@ -1,195 +0,0 @@
-from __future__ import annotations
-
-import logging
-from functools import cached_property
-from typing import Any, Dict, List, Optional
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import pre_init
-from langchain_core.utils.pydantic import get_fields
-from pydantic import BaseModel
-
-logger = logging.getLogger(__name__)
-
-MAX_BATCH_SIZE_CHARS = 1000000
-MAX_BATCH_SIZE_PARTS = 90
-
-
-@deprecated(
- since="0.3.5",
- removal="1.0",
- alternative_import="langchain_gigachat.GigaChatEmbeddings",
-)
-class GigaChatEmbeddings(BaseModel, Embeddings):
- """GigaChat Embeddings models.
-
- Example:
- .. code-block:: python
- from langchain_community.embeddings.gigachat import GigaChatEmbeddings
-
- embeddings = GigaChatEmbeddings(
- credentials=..., scope=..., verify_ssl_certs=False
- )
- """
-
- base_url: Optional[str] = None
- """ Base API URL """
- auth_url: Optional[str] = None
- """ Auth URL """
- credentials: Optional[str] = None
- """ Auth Token """
- scope: Optional[str] = None
- """ Permission scope for access token """
-
- access_token: Optional[str] = None
- """ Access token for GigaChat """
-
- model: Optional[str] = None
- """Model name to use."""
- user: Optional[str] = None
- """ Username for authenticate """
- password: Optional[str] = None
- """ Password for authenticate """
-
- timeout: Optional[float] = 600
- """ Timeout for request. By default it works for long requests. """
- verify_ssl_certs: Optional[bool] = None
- """ Check certificates for all requests """
-
- ca_bundle_file: Optional[str] = None
- cert_file: Optional[str] = None
- key_file: Optional[str] = None
- key_file_password: Optional[str] = None
- # Support for connection to GigaChat through SSL certificates
-
- @cached_property
- def _client(self) -> Any:
- """Returns GigaChat API client"""
- import gigachat
-
- return gigachat.GigaChat(
- base_url=self.base_url,
- auth_url=self.auth_url,
- credentials=self.credentials,
- scope=self.scope,
- access_token=self.access_token,
- model=self.model,
- user=self.user,
- password=self.password,
- timeout=self.timeout,
- verify_ssl_certs=self.verify_ssl_certs,
- ca_bundle_file=self.ca_bundle_file,
- cert_file=self.cert_file,
- key_file=self.key_file,
- key_file_password=self.key_file_password,
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate authenticate data in environment and python package is installed."""
- try:
- import gigachat # noqa: F401
- except ImportError:
- raise ImportError(
- "Could not import gigachat python package. "
- "Please install it with `pip install gigachat`."
- )
- fields = set(get_fields(cls).keys())
- diff = set(values.keys()) - fields
- if diff:
- logger.warning(f"Extra fields {diff} in GigaChat class")
- return values
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed documents using a GigaChat embeddings models.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- result: List[List[float]] = []
- size = 0
- local_texts = []
- embed_kwargs = {}
- if self.model is not None:
- embed_kwargs["model"] = self.model
- for text in texts:
- local_texts.append(text)
- size += len(text)
- if size > MAX_BATCH_SIZE_CHARS or len(local_texts) > MAX_BATCH_SIZE_PARTS:
- for embedding in self._client.embeddings(
- texts=local_texts, **embed_kwargs
- ).data:
- result.append(embedding.embedding)
- size = 0
- local_texts = []
- # Call for last iteration
- if local_texts:
- for embedding in self._client.embeddings(
- texts=local_texts, **embed_kwargs
- ).data:
- result.append(embedding.embedding)
-
- return result
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed documents using a GigaChat embeddings models.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- result: List[List[float]] = []
- size = 0
- local_texts = []
- embed_kwargs = {}
- if self.model is not None:
- embed_kwargs["model"] = self.model
- for text in texts:
- local_texts.append(text)
- size += len(text)
- if size > MAX_BATCH_SIZE_CHARS or len(local_texts) > MAX_BATCH_SIZE_PARTS:
- embeddings = await self._client.aembeddings(
- texts=local_texts, **embed_kwargs
- )
- for embedding in embeddings.data:
- result.append(embedding.embedding)
- size = 0
- local_texts = []
- # Call for last iteration
- if local_texts:
- embeddings = await self._client.aembeddings(
- texts=local_texts, **embed_kwargs
- )
- for embedding in embeddings.data:
- result.append(embedding.embedding)
-
- return result
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using a GigaChat embeddings models.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self.embed_documents(texts=[text])[0]
-
- async def aembed_query(self, text: str) -> List[float]:
- """Embed a query using a GigaChat embeddings models.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- docs = await self.aembed_documents(texts=[text])
- return docs[0]
diff --git a/libs/community/langchain_community/embeddings/google_palm.py b/libs/community/langchain_community/embeddings/google_palm.py
deleted file mode 100644
index d058bc46ad..0000000000
--- a/libs/community/langchain_community/embeddings/google_palm.py
+++ /dev/null
@@ -1,103 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Callable, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict
-from tenacity import (
- before_sleep_log,
- retry,
- retry_if_exception_type,
- stop_after_attempt,
- wait_exponential,
-)
-
-logger = logging.getLogger(__name__)
-
-
-def _create_retry_decorator() -> Callable[[Any], Any]:
- """Returns a tenacity retry decorator, preconfigured to handle PaLM exceptions"""
- import google.api_core.exceptions
-
- multiplier = 2
- min_seconds = 1
- max_seconds = 60
- max_retries = 10
-
- return retry(
- reraise=True,
- stop=stop_after_attempt(max_retries),
- wait=wait_exponential(multiplier=multiplier, min=min_seconds, max=max_seconds),
- retry=(
- retry_if_exception_type(google.api_core.exceptions.ResourceExhausted)
- | retry_if_exception_type(google.api_core.exceptions.ServiceUnavailable)
- | retry_if_exception_type(google.api_core.exceptions.GoogleAPIError)
- ),
- before_sleep=before_sleep_log(logger, logging.WARNING),
- )
-
-
-def embed_with_retry(
- embeddings: GooglePalmEmbeddings, *args: Any, **kwargs: Any
-) -> Any:
- """Use tenacity to retry the completion call."""
- retry_decorator = _create_retry_decorator()
-
- @retry_decorator
- def _embed_with_retry(*args: Any, **kwargs: Any) -> Any:
- return embeddings.client.generate_embeddings(*args, **kwargs)
-
- return _embed_with_retry(*args, **kwargs)
-
-
-class GooglePalmEmbeddings(BaseModel, Embeddings):
- """Google's PaLM Embeddings APIs."""
-
- client: Any
- google_api_key: Optional[str]
- model_name: str = "models/embedding-gecko-001"
- """Model name to use."""
- show_progress_bar: bool = False
- """Whether to show a tqdm progress bar. Must have `tqdm` installed."""
-
- model_config = ConfigDict(protected_namespaces=())
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate api key, python package exists."""
- google_api_key = get_from_dict_or_env(
- values, "google_api_key", "GOOGLE_API_KEY"
- )
- try:
- import google.generativeai as genai
-
- genai.configure(api_key=google_api_key)
- except ImportError:
- raise ImportError("Could not import google.generativeai python package.")
-
- values["client"] = genai
-
- return values
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- if self.show_progress_bar:
- try:
- from tqdm import tqdm
-
- iter_ = tqdm(texts, desc="GooglePalmEmbeddings")
- except ImportError:
- logger.warning(
- "Unable to show progress bar because tqdm could not be imported. "
- "Please install with `pip install tqdm`."
- )
- iter_ = texts
- else:
- iter_ = texts
- return [self.embed_query(text) for text in iter_]
-
- def embed_query(self, text: str) -> List[float]:
- """Embed query text."""
- embedding = embed_with_retry(self, self.model_name, text)
- return embedding["embedding"]
diff --git a/libs/community/langchain_community/embeddings/gpt4all.py b/libs/community/langchain_community/embeddings/gpt4all.py
deleted file mode 100644
index 5183cbb08b..0000000000
--- a/libs/community/langchain_community/embeddings/gpt4all.py
+++ /dev/null
@@ -1,76 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict, model_validator
-
-
-class GPT4AllEmbeddings(BaseModel, Embeddings):
- """GPT4All embedding models.
-
- To use, you should have the gpt4all python package installed
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import GPT4AllEmbeddings
-
- model_name = "all-MiniLM-L6-v2.gguf2.f16.gguf"
- gpt4all_kwargs = {'allow_download': 'True'}
- embeddings = GPT4AllEmbeddings(
- model_name=model_name,
- gpt4all_kwargs=gpt4all_kwargs
- )
- """
-
- model_name: Optional[str] = None
- n_threads: Optional[int] = None
- device: Optional[str] = "cpu"
- gpt4all_kwargs: Optional[dict] = {}
- client: Any #: :meta private:
-
- model_config = ConfigDict(protected_namespaces=())
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that GPT4All library is installed."""
- try:
- from gpt4all import Embed4All
-
- values["client"] = Embed4All(
- model_name=values.get("model_name"),
- n_threads=values.get("n_threads"),
- device=values.get("device"),
- **(values.get("gpt4all_kwargs") or {}),
- )
- except ImportError:
- raise ImportError(
- "Could not import gpt4all library. "
- "Please install the gpt4all library to "
- "use this embedding model: pip install gpt4all"
- )
- return values
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed a list of documents using GPT4All.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
-
- embeddings = [self.client.embed(text) for text in texts]
- return [list(map(float, e)) for e in embeddings]
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using GPT4All.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self.embed_documents([text])[0]
diff --git a/libs/community/langchain_community/embeddings/gradient_ai.py b/libs/community/langchain_community/embeddings/gradient_ai.py
deleted file mode 100644
index f33c80a6ec..0000000000
--- a/libs/community/langchain_community/embeddings/gradient_ai.py
+++ /dev/null
@@ -1,173 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import get_from_dict_or_env
-from packaging.version import parse
-from pydantic import BaseModel, ConfigDict, model_validator
-from typing_extensions import Self
-
-__all__ = ["GradientEmbeddings"]
-
-
-class GradientEmbeddings(BaseModel, Embeddings):
- """Gradient.ai Embedding models.
-
- GradientLLM is a class to interact with Embedding Models on gradient.ai
-
- To use, set the environment variable ``GRADIENT_ACCESS_TOKEN`` with your
- API token and ``GRADIENT_WORKSPACE_ID`` for your gradient workspace,
- or alternatively provide them as keywords to the constructor of this class.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import GradientEmbeddings
- GradientEmbeddings(
- model="bge-large",
- gradient_workspace_id="12345614fc0_workspace",
- gradient_access_token="gradientai-access_token",
- )
- """
-
- model: str
- "Underlying gradient.ai model id."
-
- gradient_workspace_id: Optional[str] = None
- "Underlying gradient.ai workspace_id."
-
- gradient_access_token: Optional[str] = None
- """gradient.ai API Token, which can be generated by going to
- https://auth.gradient.ai/select-workspace
- and selecting "Access tokens" under the profile drop-down.
- """
-
- gradient_api_url: str = "https://api.gradient.ai/api"
- """Endpoint URL to use."""
-
- query_prompt_for_retrieval: Optional[str] = None
- """Query pre-prompt"""
-
- client: Any = None #: :meta private:
- """Gradient client."""
-
- # LLM call kwargs
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and python package exists in environment."""
-
- values["gradient_access_token"] = get_from_dict_or_env(
- values, "gradient_access_token", "GRADIENT_ACCESS_TOKEN"
- )
- values["gradient_workspace_id"] = get_from_dict_or_env(
- values, "gradient_workspace_id", "GRADIENT_WORKSPACE_ID"
- )
-
- values["gradient_api_url"] = get_from_dict_or_env(
- values,
- "gradient_api_url",
- "GRADIENT_API_URL",
- default="https://api.gradient.ai/api",
- )
- return values
-
- @model_validator(mode="after")
- def post_init(self) -> Self:
- try:
- import gradientai
- except ImportError:
- raise ImportError(
- 'GradientEmbeddings requires `pip install -U "gradientai>=1.4.0"`.'
- )
-
- if parse(gradientai.__version__) < parse("1.4.0"):
- raise ImportError(
- 'GradientEmbeddings requires `pip install -U "gradientai>=1.4.0"`.'
- )
-
- gradient = gradientai.Gradient(
- access_token=self.gradient_access_token,
- workspace_id=self.gradient_workspace_id,
- host=self.gradient_api_url,
- )
- self.client = gradient.get_embeddings_model(slug=self.model)
- return self
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Call out to Gradient's embedding endpoint.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- inputs = [{"input": text} for text in texts]
-
- result = self.client.embed(inputs=inputs).embeddings
-
- return [e.embedding for e in result]
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- """Async call out to Gradient's embedding endpoint.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- inputs = [{"input": text} for text in texts]
-
- result = (await self.client.aembed(inputs=inputs)).embeddings
-
- return [e.embedding for e in result]
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to Gradient's embedding endpoint.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- query = (
- f"{self.query_prompt_for_retrieval} {text}"
- if self.query_prompt_for_retrieval
- else text
- )
- return self.embed_documents([query])[0]
-
- async def aembed_query(self, text: str) -> List[float]:
- """Async call out to Gradient's embedding endpoint.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- query = (
- f"{self.query_prompt_for_retrieval} {text}"
- if self.query_prompt_for_retrieval
- else text
- )
- embeddings = await self.aembed_documents([query])
- return embeddings[0]
-
-
-class TinyAsyncGradientEmbeddingClient: #: :meta private:
- """Deprecated, TinyAsyncGradientEmbeddingClient was removed.
-
- This class is just for backwards compatibility with older versions
- of langchain_community.
- It might be entirely removed in the future.
- """
-
- def __init__(self, *args, **kwargs) -> None: # type: ignore[no-untyped-def]
- raise ValueError("Deprecated,TinyAsyncGradientEmbeddingClient was removed.")
diff --git a/libs/community/langchain_community/embeddings/huggingface.py b/libs/community/langchain_community/embeddings/huggingface.py
deleted file mode 100644
index 5810bbc920..0000000000
--- a/libs/community/langchain_community/embeddings/huggingface.py
+++ /dev/null
@@ -1,483 +0,0 @@
-import warnings
-from typing import Any, Dict, List, Optional
-
-import requests
-from langchain_core._api import deprecated, warn_deprecated
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict, Field, SecretStr
-
-DEFAULT_MODEL_NAME = "sentence-transformers/all-mpnet-base-v2"
-DEFAULT_INSTRUCT_MODEL = "hkunlp/instructor-large"
-DEFAULT_BGE_MODEL = "BAAI/bge-large-en"
-DEFAULT_EMBED_INSTRUCTION = "Represent the document for retrieval: "
-DEFAULT_QUERY_INSTRUCTION = (
- "Represent the question for retrieving supporting documents: "
-)
-DEFAULT_QUERY_BGE_INSTRUCTION_EN = (
- "Represent this question for searching relevant passages: "
-)
-DEFAULT_QUERY_BGE_INSTRUCTION_ZH = "为这个句子生成表示以用于检索相关文章:"
-
-
-@deprecated(
- since="0.2.2",
- removal="1.0",
- alternative_import="langchain_huggingface.HuggingFaceEmbeddings",
-)
-class HuggingFaceEmbeddings(BaseModel, Embeddings):
- """HuggingFace sentence_transformers embedding models.
-
- To use, you should have the ``sentence_transformers`` python package installed.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import HuggingFaceEmbeddings
-
- model_name = "sentence-transformers/all-mpnet-base-v2"
- model_kwargs = {'device': 'cpu'}
- encode_kwargs = {'normalize_embeddings': False}
- hf = HuggingFaceEmbeddings(
- model_name=model_name,
- model_kwargs=model_kwargs,
- encode_kwargs=encode_kwargs
- )
- """
-
- client: Any = None #: :meta private:
- model_name: str = DEFAULT_MODEL_NAME
- """Model name to use."""
- cache_folder: Optional[str] = None
- """Path to store models.
- Can be also set by SENTENCE_TRANSFORMERS_HOME environment variable."""
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Keyword arguments to pass to the Sentence Transformer model, such as `device`,
- `prompts`, `default_prompt_name`, `revision`, `trust_remote_code`, or `token`.
- See also the Sentence Transformer documentation: https://sbert.net/docs/package_reference/SentenceTransformer.html#sentence_transformers.SentenceTransformer"""
- encode_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Keyword arguments to pass when calling the `encode` method of the Sentence
- Transformer model, such as `prompt_name`, `prompt`, `batch_size`, `precision`,
- `normalize_embeddings`, and more.
- See also the Sentence Transformer documentation: https://sbert.net/docs/package_reference/SentenceTransformer.html#sentence_transformers.SentenceTransformer.encode"""
- multi_process: bool = False
- """Run encode() on multiple GPUs."""
- show_progress: bool = False
- """Whether to show a progress bar."""
-
- def __init__(self, **kwargs: Any):
- """Initialize the sentence_transformer."""
- super().__init__(**kwargs)
-
- if "model_name" not in kwargs:
- since = "0.2.16"
- removal = "0.4.0"
- warn_deprecated(
- since=since,
- removal=removal,
- message=f"Default values for {self.__class__.__name__}.model_name"
- + f" were deprecated in LangChain {since} and will be removed in"
- + f" {removal}. Explicitly pass a model_name to the"
- + f" {self.__class__.__name__} constructor instead.",
- )
-
- try:
- import sentence_transformers
-
- except ImportError as exc:
- raise ImportError(
- "Could not import sentence_transformers python package. "
- "Please install it with `pip install sentence-transformers`."
- ) from exc
-
- self.client = sentence_transformers.SentenceTransformer(
- self.model_name, cache_folder=self.cache_folder, **self.model_kwargs
- )
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Compute doc embeddings using a HuggingFace transformer model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- import sentence_transformers
-
- texts = list(map(lambda x: x.replace("\n", " "), texts))
- if self.multi_process:
- pool = self.client.start_multi_process_pool()
- embeddings = self.client.encode_multi_process(texts, pool)
- sentence_transformers.SentenceTransformer.stop_multi_process_pool(pool)
- else:
- embeddings = self.client.encode(
- texts, show_progress_bar=self.show_progress, **self.encode_kwargs
- )
-
- return embeddings.tolist()
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using a HuggingFace transformer model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self.embed_documents([text])[0]
-
-
-class HuggingFaceInstructEmbeddings(BaseModel, Embeddings):
- """Wrapper around sentence_transformers embedding models.
-
- To use, you should have the ``sentence_transformers``
- and ``InstructorEmbedding`` python packages installed.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import HuggingFaceInstructEmbeddings
-
- model_name = "hkunlp/instructor-large"
- model_kwargs = {'device': 'cpu'}
- encode_kwargs = {'normalize_embeddings': True}
- hf = HuggingFaceInstructEmbeddings(
- model_name=model_name,
- model_kwargs=model_kwargs,
- encode_kwargs=encode_kwargs
- )
- """
-
- client: Any = None #: :meta private:
- model_name: str = DEFAULT_INSTRUCT_MODEL
- """Model name to use."""
- cache_folder: Optional[str] = None
- """Path to store models.
- Can be also set by SENTENCE_TRANSFORMERS_HOME environment variable."""
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Keyword arguments to pass to the model."""
- encode_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Keyword arguments to pass when calling the `encode` method of the model."""
- embed_instruction: str = DEFAULT_EMBED_INSTRUCTION
- """Instruction to use for embedding documents."""
- query_instruction: str = DEFAULT_QUERY_INSTRUCTION
- """Instruction to use for embedding query."""
- show_progress: bool = False
- """Whether to show a progress bar."""
-
- def __init__(self, **kwargs: Any):
- """Initialize the sentence_transformer."""
- super().__init__(**kwargs)
-
- if "model_name" not in kwargs:
- since = "0.2.16"
- removal = "0.4.0"
- warn_deprecated(
- since=since,
- removal=removal,
- message=f"Default values for {self.__class__.__name__}.model_name"
- + f" were deprecated in LangChain {since} and will be removed in"
- + f" {removal}. Explicitly pass a model_name to the"
- + f" {self.__class__.__name__} constructor instead.",
- )
-
- try:
- from InstructorEmbedding import INSTRUCTOR
-
- self.client = INSTRUCTOR(
- self.model_name, cache_folder=self.cache_folder, **self.model_kwargs
- )
- except ImportError as e:
- raise ImportError("Dependencies for InstructorEmbedding not found.") from e
-
- if "show_progress_bar" in self.encode_kwargs:
- warn_deprecated(
- since="0.2.5",
- removal="1.0",
- name="encode_kwargs['show_progress_bar']",
- alternative=f"the show_progress method on {self.__class__.__name__}",
- )
- if self.show_progress:
- warnings.warn(
- "Both encode_kwargs['show_progress_bar'] and show_progress are set;"
- "encode_kwargs['show_progress_bar'] takes precedence"
- )
- self.show_progress = self.encode_kwargs.pop("show_progress_bar")
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Compute doc embeddings using a HuggingFace instruct model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- instruction_pairs = [[self.embed_instruction, text] for text in texts]
- embeddings = self.client.encode(
- instruction_pairs,
- show_progress_bar=self.show_progress,
- **self.encode_kwargs,
- )
- return embeddings.tolist()
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using a HuggingFace instruct model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- instruction_pair = [self.query_instruction, text]
- embedding = self.client.encode(
- [instruction_pair],
- show_progress_bar=self.show_progress,
- **self.encode_kwargs,
- )[0]
- return embedding.tolist()
-
-
-@deprecated(
- since="0.2.2",
- removal="1.0",
- alternative_import="langchain_huggingface.HuggingFaceEmbeddings",
-)
-class HuggingFaceBgeEmbeddings(BaseModel, Embeddings):
- """HuggingFace sentence_transformers embedding models.
-
- To use, you should have the ``sentence_transformers`` python package installed.
- To use Nomic, make sure the version of ``sentence_transformers`` >= 2.3.0.
-
- Bge Example:
- .. code-block:: python
-
- from langchain_community.embeddings import HuggingFaceBgeEmbeddings
-
- model_name = "BAAI/bge-large-en-v1.5"
- model_kwargs = {'device': 'cpu'}
- encode_kwargs = {'normalize_embeddings': True}
- hf = HuggingFaceBgeEmbeddings(
- model_name=model_name,
- model_kwargs=model_kwargs,
- encode_kwargs=encode_kwargs
- )
- Nomic Example:
- .. code-block:: python
-
- from langchain_community.embeddings import HuggingFaceBgeEmbeddings
-
- model_name = "nomic-ai/nomic-embed-text-v1"
- model_kwargs = {
- 'device': 'cpu',
- 'trust_remote_code':True
- }
- encode_kwargs = {'normalize_embeddings': True}
- hf = HuggingFaceBgeEmbeddings(
- model_name=model_name,
- model_kwargs=model_kwargs,
- encode_kwargs=encode_kwargs,
- query_instruction = "search_query:",
- embed_instruction = "search_document:"
- )
- """
-
- client: Any = None #: :meta private:
- model_name: str = DEFAULT_BGE_MODEL
- """Model name to use."""
- cache_folder: Optional[str] = None
- """Path to store models.
- Can be also set by SENTENCE_TRANSFORMERS_HOME environment variable."""
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Keyword arguments to pass to the model."""
- encode_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Keyword arguments to pass when calling the `encode` method of the model."""
- query_instruction: str = DEFAULT_QUERY_BGE_INSTRUCTION_EN
- """Instruction to use for embedding query."""
- embed_instruction: str = ""
- """Instruction to use for embedding document."""
- show_progress: bool = False
- """Whether to show a progress bar."""
-
- def __init__(self, **kwargs: Any):
- """Initialize the sentence_transformer."""
- super().__init__(**kwargs)
-
- if "model_name" not in kwargs:
- since = "0.2.5"
- removal = "0.4.0"
- warn_deprecated(
- since=since,
- removal=removal,
- message=f"Default values for {self.__class__.__name__}.model_name"
- + f" were deprecated in LangChain {since} and will be removed in"
- + f" {removal}. Explicitly pass a model_name to the"
- + f" {self.__class__.__name__} constructor instead.",
- )
-
- try:
- import sentence_transformers
-
- except ImportError as exc:
- raise ImportError(
- "Could not import sentence_transformers python package. "
- "Please install it with `pip install sentence-transformers`."
- ) from exc
- extra_model_kwargs = [
- "torch_dtype",
- "attn_implementation",
- "provider",
- "file_name",
- "export",
- ]
- extra_model_kwargs_dict = {
- k: self.model_kwargs.pop(k)
- for k in extra_model_kwargs
- if k in self.model_kwargs
- }
- self.client = sentence_transformers.SentenceTransformer(
- self.model_name,
- cache_folder=self.cache_folder,
- **self.model_kwargs,
- model_kwargs=extra_model_kwargs_dict,
- )
-
- if "-zh" in self.model_name:
- self.query_instruction = DEFAULT_QUERY_BGE_INSTRUCTION_ZH
-
- if "show_progress_bar" in self.encode_kwargs:
- warn_deprecated(
- since="0.2.5",
- removal="1.0",
- name="encode_kwargs['show_progress_bar']",
- alternative=f"the show_progress method on {self.__class__.__name__}",
- )
- if self.show_progress:
- warnings.warn(
- "Both encode_kwargs['show_progress_bar'] and show_progress are set;"
- "encode_kwargs['show_progress_bar'] takes precedence"
- )
- self.show_progress = self.encode_kwargs.pop("show_progress_bar")
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Compute doc embeddings using a HuggingFace transformer model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- texts = [self.embed_instruction + t.replace("\n", " ") for t in texts]
- embeddings = self.client.encode(
- texts, show_progress_bar=self.show_progress, **self.encode_kwargs
- )
- return embeddings.tolist()
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using a HuggingFace transformer model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- text = text.replace("\n", " ")
- embedding = self.client.encode(
- self.query_instruction + text,
- show_progress_bar=self.show_progress,
- **self.encode_kwargs,
- )
- return embedding.tolist()
-
-
-class HuggingFaceInferenceAPIEmbeddings(BaseModel, Embeddings):
- """Embed texts using the HuggingFace API.
-
- Requires a HuggingFace Inference API key and a model name.
- """
-
- api_key: SecretStr
- """Your API key for the HuggingFace Inference API."""
- model_name: str = "sentence-transformers/all-MiniLM-L6-v2"
- """The name of the model to use for text embeddings."""
- api_url: Optional[str] = None
- """Custom inference endpoint url. None for using default public url."""
- additional_headers: Dict[str, str] = {}
- """Pass additional headers to the requests library if needed."""
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- @property
- def _api_url(self) -> str:
- return self.api_url or self._default_api_url
-
- @property
- def _default_api_url(self) -> str:
- return (
- "https://api-inference.huggingface.co"
- "/pipeline"
- "/feature-extraction"
- f"/{self.model_name}"
- )
-
- @property
- def _headers(self) -> dict:
- return {
- "Authorization": f"Bearer {self.api_key.get_secret_value()}",
- **self.additional_headers,
- }
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Get the embeddings for a list of texts.
-
- Args:
- texts (Documents): A list of texts to get embeddings for.
-
- Returns:
- Embedded texts as List[List[float]], where each inner List[float]
- corresponds to a single input text.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import (
- HuggingFaceInferenceAPIEmbeddings,
- )
-
- hf_embeddings = HuggingFaceInferenceAPIEmbeddings(
- api_key="your_api_key",
- model_name="sentence-transformers/all-MiniLM-l6-v2"
- )
- texts = ["Hello, world!", "How are you?"]
- hf_embeddings.embed_documents(texts)
- """ # noqa: E501
- response = requests.post(
- self._api_url,
- headers=self._headers,
- json={
- "inputs": texts,
- "options": {"wait_for_model": True, "use_cache": True},
- },
- )
- return response.json()
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using a HuggingFace transformer model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self.embed_documents([text])[0]
diff --git a/libs/community/langchain_community/embeddings/huggingface_hub.py b/libs/community/langchain_community/embeddings/huggingface_hub.py
deleted file mode 100644
index b1a1fac372..0000000000
--- a/libs/community/langchain_community/embeddings/huggingface_hub.py
+++ /dev/null
@@ -1,159 +0,0 @@
-import json
-from typing import Any, Dict, List, Optional
-
-from langchain_core._api import deprecated
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import BaseModel, ConfigDict, model_validator
-from typing_extensions import Self
-
-DEFAULT_MODEL = "sentence-transformers/all-mpnet-base-v2"
-VALID_TASKS = ("feature-extraction",)
-
-
-@deprecated(
- since="0.2.2",
- removal="1.0",
- alternative_import="langchain_huggingface.HuggingFaceEndpointEmbeddings",
-)
-class HuggingFaceHubEmbeddings(BaseModel, Embeddings):
- """HuggingFaceHub embedding models.
-
- To use, you should have the ``huggingface_hub`` python package installed, and the
- environment variable ``HUGGINGFACEHUB_API_TOKEN`` set with your API token, or pass
- it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import HuggingFaceHubEmbeddings
- model = "sentence-transformers/all-mpnet-base-v2"
- hf = HuggingFaceHubEmbeddings(
- model=model,
- task="feature-extraction",
- huggingfacehub_api_token="my-api-key",
- )
- """
-
- client: Any = None #: :meta private:
- async_client: Any = None #: :meta private:
- model: Optional[str] = None
- """Model name to use."""
- repo_id: Optional[str] = None
- """Huggingfacehub repository id, for backward compatibility."""
- task: Optional[str] = "feature-extraction"
- """Task to call the model with."""
- model_kwargs: Optional[dict] = None
- """Keyword arguments to pass to the model."""
-
- huggingfacehub_api_token: Optional[str] = None
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and python package exists in environment."""
- huggingfacehub_api_token = get_from_dict_or_env(
- values, "huggingfacehub_api_token", "HUGGINGFACEHUB_API_TOKEN"
- )
-
- try:
- from huggingface_hub import AsyncInferenceClient, InferenceClient
-
- if values.get("model"):
- values["repo_id"] = values["model"]
- elif values.get("repo_id"):
- values["model"] = values["repo_id"]
- else:
- values["model"] = DEFAULT_MODEL
- values["repo_id"] = DEFAULT_MODEL
-
- client = InferenceClient(
- model=values["model"],
- token=huggingfacehub_api_token,
- )
-
- async_client = AsyncInferenceClient(
- model=values["model"],
- token=huggingfacehub_api_token,
- )
-
- values["client"] = client
- values["async_client"] = async_client
-
- except ImportError:
- raise ImportError(
- "Could not import huggingface_hub python package. "
- "Please install it with `pip install huggingface_hub`."
- )
- return values
-
- @model_validator(mode="after")
- def post_init(self) -> Self:
- """Post init validation for the class."""
- if self.task not in VALID_TASKS:
- raise ValueError(
- f"Got invalid task {self.task}, "
- f"currently only {VALID_TASKS} are supported"
- )
- return self
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Call out to HuggingFaceHub's embedding endpoint for embedding search docs.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- # replace newlines, which can negatively affect performance.
- texts = [text.replace("\n", " ") for text in texts]
- _model_kwargs = self.model_kwargs or {}
- # api doc: https://huggingface.github.io/text-embeddings-inference/#/Text%20Embeddings%20Inference/embed
- responses = self.client.post(
- json={"inputs": texts, **_model_kwargs}, task=self.task
- )
- return json.loads(responses.decode())
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- """Async Call to HuggingFaceHub's embedding endpoint for embedding search docs.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- # replace newlines, which can negatively affect performance.
- texts = [text.replace("\n", " ") for text in texts]
- _model_kwargs = self.model_kwargs or {}
- responses = await self.async_client.post(
- json={"inputs": texts, "parameters": _model_kwargs}, task=self.task
- )
- return json.loads(responses.decode())
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to HuggingFaceHub's embedding endpoint for embedding query text.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- response = self.embed_documents([text])[0]
- return response
-
- async def aembed_query(self, text: str) -> List[float]:
- """Async Call to HuggingFaceHub's embedding endpoint for embedding query text.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- response = (await self.aembed_documents([text]))[0]
- return response
diff --git a/libs/community/langchain_community/embeddings/hunyuan.py b/libs/community/langchain_community/embeddings/hunyuan.py
deleted file mode 100644
index 1d0570a0ae..0000000000
--- a/libs/community/langchain_community/embeddings/hunyuan.py
+++ /dev/null
@@ -1,124 +0,0 @@
-import json
-from typing import Any, Dict, List, Literal, Optional, Type
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.runnables.config import run_in_executor
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env
-from pydantic import BaseModel, Field, SecretStr, model_validator
-
-
-class HunyuanEmbeddings(Embeddings, BaseModel):
- """Tencent Hunyuan embedding models API by Tencent.
-
- For more information, see https://cloud.tencent.com/document/product/1729
- """
-
- hunyuan_secret_id: Optional[SecretStr] = Field(alias="secret_id", default=None)
- """Hunyuan Secret ID"""
- hunyuan_secret_key: Optional[SecretStr] = Field(alias="secret_key", default=None)
- """Hunyuan Secret Key"""
- region: Literal["ap-guangzhou", "ap-beijing"] = "ap-guangzhou"
- """The region of hunyuan service."""
- embedding_ctx_length: int = 1024
- """The max embedding context length of hunyuan embedding (defaults to 1024)."""
- show_progress_bar: bool = False
- """Show progress bar when embedding. Default is False."""
-
- client: Any = Field(default=None, exclude=True)
- """The tencentcloud client."""
- request_cls: Optional[Type] = Field(default=None, exclude=True)
- """The request class of tencentcloud sdk."""
-
- @model_validator(mode="before")
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- values["hunyuan_secret_id"] = convert_to_secret_str(
- get_from_dict_or_env(
- values,
- "hunyuan_secret_id",
- "HUNYUAN_SECRET_ID",
- )
- )
- values["hunyuan_secret_key"] = convert_to_secret_str(
- get_from_dict_or_env(
- values,
- "hunyuan_secret_key",
- "HUNYUAN_SECRET_KEY",
- )
- )
-
- try:
- from tencentcloud.common.credential import Credential
- from tencentcloud.common.profile.client_profile import ClientProfile
- from tencentcloud.hunyuan.v20230901.hunyuan_client import HunyuanClient
- from tencentcloud.hunyuan.v20230901.models import GetEmbeddingRequest
- except ImportError:
- raise ImportError(
- "Could not import tencentcloud sdk python package. Please install it "
- 'with `pip install "tencentcloud-sdk-python>=3.0.1139"`.'
- )
-
- client_profile = ClientProfile()
- client_profile.httpProfile.pre_conn_pool_size = 3
-
- credential = Credential(
- values["hunyuan_secret_id"].get_secret_value(),
- values["hunyuan_secret_key"].get_secret_value(),
- )
-
- values["request_cls"] = GetEmbeddingRequest
-
- values["client"] = HunyuanClient(credential, values["region"], client_profile)
- return values
-
- def _embed_text(self, text: str) -> List[float]:
- if self.request_cls is None:
- raise AssertionError("Request class is not initialized.")
- request = self.request_cls()
- request.Input = text
-
- response = self.client.GetEmbedding(request)
-
- _response: Dict[str, Any] = json.loads(response.to_json_string())
-
- data: Optional[List[Dict[str, Any]]] = _response.get("Data")
- if not data:
- raise RuntimeError("Occur hunyuan embedding error: Data is empty")
-
- embedding = data[0].get("Embedding")
- if not embedding:
- raise RuntimeError("Occur hunyuan embedding error: Embedding is empty")
-
- return embedding
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed search docs."""
- embeddings = []
- if self.show_progress_bar:
- try:
- from tqdm import tqdm
- except ImportError as e:
- raise ImportError(
- "Package tqdm must be installed if show_progress_bar=True. "
- "Please install with 'pip install tqdm' or set "
- "show_progress_bar=False."
- ) from e
- _iter = tqdm(iterable=texts, desc="Hunyuan Embedding")
- else:
- _iter = texts
- for text in _iter:
- embeddings.append(self.embed_query(text))
-
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Embed query text."""
- return self._embed_text(text)
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- """Asynchronous Embed search docs."""
- return await run_in_executor(None, self.embed_documents, texts)
-
- async def aembed_query(self, text: str) -> List[float]:
- """Asynchronous Embed query text."""
- return await run_in_executor(None, self.embed_query, text)
diff --git a/libs/community/langchain_community/embeddings/infinity.py b/libs/community/langchain_community/embeddings/infinity.py
deleted file mode 100644
index cc41250b54..0000000000
--- a/libs/community/langchain_community/embeddings/infinity.py
+++ /dev/null
@@ -1,324 +0,0 @@
-"""written under MIT Licence, Michael Feil 2023."""
-
-import asyncio
-from concurrent.futures import ThreadPoolExecutor
-from typing import Any, Callable, Dict, List, Optional, Tuple
-
-import aiohttp
-import numpy as np
-import requests
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import BaseModel, ConfigDict, model_validator
-
-__all__ = ["InfinityEmbeddings"]
-
-
-class InfinityEmbeddings(BaseModel, Embeddings):
- """Self-hosted embedding models for `infinity` package.
-
- See https://github.com/michaelfeil/infinity
- This also works for text-embeddings-inference and other
- self-hosted openai-compatible servers.
-
- Infinity is a package to interact with Embedding Models on https://github.com/michaelfeil/infinity
-
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import InfinityEmbeddings
- InfinityEmbeddings(
- model="BAAI/bge-small",
- infinity_api_url="http://localhost:7997",
- )
- """
-
- model: str
- "Underlying Infinity model id."
-
- infinity_api_url: str = "http://localhost:7997"
- """Endpoint URL to use."""
-
- client: Any = None #: :meta private:
- """Infinity client."""
-
- # LLM call kwargs
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and python package exists in environment."""
-
- values["infinity_api_url"] = get_from_dict_or_env(
- values, "infinity_api_url", "INFINITY_API_URL"
- )
-
- values["client"] = TinyAsyncOpenAIInfinityEmbeddingClient(
- host=values["infinity_api_url"],
- )
- return values
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Call out to Infinity's embedding endpoint.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- embeddings = self.client.embed(
- model=self.model,
- texts=texts,
- )
- return embeddings
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- """Async call out to Infinity's embedding endpoint.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- embeddings = await self.client.aembed(
- model=self.model,
- texts=texts,
- )
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to Infinity's embedding endpoint.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self.embed_documents([text])[0]
-
- async def aembed_query(self, text: str) -> List[float]:
- """Async call out to Infinity's embedding endpoint.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- embeddings = await self.aembed_documents([text])
- return embeddings[0]
-
-
-class TinyAsyncOpenAIInfinityEmbeddingClient: #: :meta private:
- """Helper tool to embed Infinity.
-
- It is not a part of Langchain's stable API,
- direct use discouraged.
-
- Example:
- .. code-block:: python
-
-
- mini_client = TinyAsyncInfinityEmbeddingClient(
- )
- embeds = mini_client.embed(
- model="BAAI/bge-small",
- text=["doc1", "doc2"]
- )
- # or
- embeds = await mini_client.aembed(
- model="BAAI/bge-small",
- text=["doc1", "doc2"]
- )
-
- """
-
- def __init__(
- self,
- host: str = "http://localhost:7797/v1",
- aiosession: Optional[aiohttp.ClientSession] = None,
- ) -> None:
- self.host = host
- self.aiosession = aiosession
-
- if self.host is None or len(self.host) < 3:
- raise ValueError(" param `host` must be set to a valid url")
- self._batch_size = 128
-
- @staticmethod
- def _permute(
- texts: List[str], sorter: Callable = len
- ) -> Tuple[List[str], Callable]:
- """Sort texts in ascending order, and
- delivers a lambda expr, which can sort a same length list
- https://github.com/UKPLab/sentence-transformers/blob/
- c5f93f70eca933c78695c5bc686ceda59651ae3b/sentence_transformers/SentenceTransformer.py#L156
-
- Args:
- texts (List[str]): _description_
- sorter (Callable, optional): _description_. Defaults to len.
-
- Returns:
- Tuple[List[str], Callable]: _description_
-
- Example:
- ```
- texts = ["one","three","four"]
- perm_texts, undo = self._permute(texts)
- texts == undo(perm_texts)
- ```
- """
-
- if len(texts) == 1:
- # special case query
- return texts, lambda t: t
- length_sorted_idx = np.argsort([-sorter(sen) for sen in texts])
- texts_sorted = [texts[idx] for idx in length_sorted_idx]
-
- return texts_sorted, lambda unsorted_embeddings: [ # E731
- unsorted_embeddings[idx] for idx in np.argsort(length_sorted_idx)
- ]
-
- def _batch(self, texts: List[str]) -> List[List[str]]:
- """
- splits Lists of text parts into batches of size max `self._batch_size`
- When encoding vector database,
-
- Args:
- texts (List[str]): List of sentences
- self._batch_size (int, optional): max batch size of one request.
-
- Returns:
- List[List[str]]: Batches of List of sentences
- """
- if len(texts) == 1:
- # special case query
- return [texts]
- batches = []
- for start_index in range(0, len(texts), self._batch_size):
- batches.append(texts[start_index : start_index + self._batch_size])
- return batches
-
- @staticmethod
- def _unbatch(batch_of_texts: List[List[Any]]) -> List[Any]:
- if len(batch_of_texts) == 1 and len(batch_of_texts[0]) == 1:
- # special case query
- return batch_of_texts[0]
- texts = []
- for sublist in batch_of_texts:
- texts.extend(sublist)
- return texts
-
- def _kwargs_post_request(self, model: str, texts: List[str]) -> Dict[str, Any]:
- """Build the kwargs for the Post request, used by sync
-
- Args:
- model (str): _description_
- texts (List[str]): _description_
-
- Returns:
- Dict[str, Collection[str]]: _description_
- """
- return dict(
- url=f"{self.host}/embeddings",
- headers={
- # "accept": "application/json",
- "content-type": "application/json",
- },
- json=dict(
- input=texts,
- model=model,
- ),
- )
-
- def _sync_request_embed(
- self, model: str, batch_texts: List[str]
- ) -> List[List[float]]:
- response = requests.post(
- **self._kwargs_post_request(model=model, texts=batch_texts)
- )
- if response.status_code != 200:
- raise Exception(
- f"Infinity returned an unexpected response with status "
- f"{response.status_code}: {response.text}"
- )
- return [e["embedding"] for e in response.json()["data"]]
-
- def embed(self, model: str, texts: List[str]) -> List[List[float]]:
- """call the embedding of model
-
- Args:
- model (str): to embedding model
- texts (List[str]): List of sentences to embed.
-
- Returns:
- List[List[float]]: List of vectors for each sentence
- """
- perm_texts, unpermute_func = self._permute(texts)
- perm_texts_batched = self._batch(perm_texts)
-
- # Request
- map_args = (
- self._sync_request_embed,
- [model] * len(perm_texts_batched),
- perm_texts_batched,
- )
- if len(perm_texts_batched) == 1:
- embeddings_batch_perm = list(map(*map_args))
- else:
- with ThreadPoolExecutor(32) as p:
- embeddings_batch_perm = list(p.map(*map_args))
-
- embeddings_perm = self._unbatch(embeddings_batch_perm)
- embeddings = unpermute_func(embeddings_perm)
- return embeddings
-
- async def _async_request(
- self, session: aiohttp.ClientSession, kwargs: Dict[str, Any]
- ) -> List[List[float]]:
- async with session.post(**kwargs) as response:
- if response.status != 200:
- raise Exception(
- f"Infinity returned an unexpected response with status "
- f"{response.status}: {response.text}"
- )
- embedding = (await response.json())["data"]
- return [e["embedding"] for e in embedding]
-
- async def aembed(self, model: str, texts: List[str]) -> List[List[float]]:
- """call the embedding of model, async method
-
- Args:
- model (str): to embedding model
- texts (List[str]): List of sentences to embed.
-
- Returns:
- List[List[float]]: List of vectors for each sentence
- """
- perm_texts, unpermute_func = self._permute(texts)
- perm_texts_batched = self._batch(perm_texts)
-
- # Request
- async with aiohttp.ClientSession(
- trust_env=True, connector=aiohttp.TCPConnector(limit=32)
- ) as session:
- embeddings_batch_perm = await asyncio.gather(
- *[
- self._async_request(
- session=session,
- kwargs=self._kwargs_post_request(model=model, texts=t),
- )
- for t in perm_texts_batched
- ]
- )
-
- embeddings_perm = self._unbatch(embeddings_batch_perm)
- embeddings = unpermute_func(embeddings_perm)
- return embeddings
diff --git a/libs/community/langchain_community/embeddings/infinity_local.py b/libs/community/langchain_community/embeddings/infinity_local.py
deleted file mode 100644
index 22e15b017a..0000000000
--- a/libs/community/langchain_community/embeddings/infinity_local.py
+++ /dev/null
@@ -1,157 +0,0 @@
-"""written under MIT Licence, Michael Feil 2023."""
-
-import asyncio
-from logging import getLogger
-from typing import Any, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict, model_validator
-from typing_extensions import Self
-
-__all__ = ["InfinityEmbeddingsLocal"]
-
-logger = getLogger(__name__)
-
-
-class InfinityEmbeddingsLocal(BaseModel, Embeddings):
- """Optimized Infinity embedding models.
-
- https://github.com/michaelfeil/infinity
- This class deploys a local Infinity instance to embed text.
- The class requires async usage.
-
- Infinity is a class to interact with Embedding Models on https://github.com/michaelfeil/infinity
-
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import InfinityEmbeddingsLocal
- async with InfinityEmbeddingsLocal(
- model="BAAI/bge-small-en-v1.5",
- revision=None,
- device="cpu",
- ) as embedder:
- embeddings = await engine.aembed_documents(["text1", "text2"])
- """
-
- model: str
- "Underlying model id from huggingface, e.g. BAAI/bge-small-en-v1.5"
-
- revision: Optional[str] = None
- "Model version, the commit hash from huggingface"
-
- batch_size: int = 32
- "Internal batch size for inference, e.g. 32"
-
- device: str = "auto"
- "Device to use for inference, e.g. 'cpu' or 'cuda', or 'mps'"
-
- backend: str = "torch"
- "Backend for inference, e.g. 'torch' (recommended for ROCm/Nvidia)"
- " or 'optimum' for onnx/tensorrt"
-
- model_warmup: bool = True
- "Warmup the model with the max batch size."
-
- engine: Any = None #: :meta private:
- """Infinity's AsyncEmbeddingEngine."""
-
- # LLM call kwargs
- model_config = ConfigDict(
- extra="forbid",
- protected_namespaces=(),
- )
-
- @model_validator(mode="after")
- def validate_environment(self) -> Self:
- """Validate that api key and python package exists in environment."""
-
- try:
- from infinity_emb import AsyncEmbeddingEngine
- except ImportError:
- raise ImportError(
- "Please install the "
- "`pip install 'infinity_emb[optimum,torch]>=0.0.24'` "
- "package to use the InfinityEmbeddingsLocal."
- )
- self.engine = AsyncEmbeddingEngine(
- model_name_or_path=self.model,
- device=self.device,
- revision=self.revision,
- model_warmup=self.model_warmup,
- batch_size=self.batch_size,
- engine=self.backend,
- )
- return self
-
- async def __aenter__(self) -> None:
- """start the background worker.
- recommended usage is with the async with statement.
-
- async with InfinityEmbeddingsLocal(
- model="BAAI/bge-small-en-v1.5",
- revision=None,
- device="cpu",
- ) as embedder:
- embeddings = await engine.aembed_documents(["text1", "text2"])
- """
- await self.engine.__aenter__()
-
- async def __aexit__(self, *args: Any) -> None:
- """stop the background worker,
- required to free references to the pytorch model."""
- await self.engine.__aexit__(*args)
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- """Async call out to Infinity's embedding endpoint.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- if not self.engine.running:
- logger.warning(
- "Starting Infinity engine on the fly. This is not recommended."
- "Please start the engine before using it."
- )
- async with self:
- # spawning threadpool for multithreaded encode, tokenization
- embeddings, _ = await self.engine.embed(texts)
- # stopping threadpool on exit
- logger.warning("Stopped infinity engine after usage.")
- else:
- embeddings, _ = await self.engine.embed(texts)
- return embeddings
-
- async def aembed_query(self, text: str) -> List[float]:
- """Async call out to Infinity's embedding endpoint.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- embeddings = await self.aembed_documents([text])
- return embeddings[0]
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """
- This method is async only.
- """
- logger.warning(
- "This method is async only. "
- "Please use the async version `await aembed_documents`."
- )
- return asyncio.run(self.aembed_documents(texts))
-
- def embed_query(self, text: str) -> List[float]:
- """ """
- logger.warning(
- "This method is async only."
- " Please use the async version `await aembed_query`."
- )
- return asyncio.run(self.aembed_query(text))
diff --git a/libs/community/langchain_community/embeddings/ipex_llm.py b/libs/community/langchain_community/embeddings/ipex_llm.py
deleted file mode 100644
index 8022616f22..0000000000
--- a/libs/community/langchain_community/embeddings/ipex_llm.py
+++ /dev/null
@@ -1,137 +0,0 @@
-# This file is adapted from
-# https://github.com/langchain-ai/langchain/blob/master/libs/community/langchain_community/embeddings/huggingface.py
-
-from typing import Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict, Field
-
-DEFAULT_BGE_MODEL = "BAAI/bge-small-en-v1.5"
-DEFAULT_QUERY_BGE_INSTRUCTION_EN = (
- "Represent this question for searching relevant passages: "
-)
-DEFAULT_QUERY_BGE_INSTRUCTION_ZH = "为这个句子生成表示以用于检索相关文章:"
-
-
-class IpexLLMBgeEmbeddings(BaseModel, Embeddings):
- """Wrapper around the BGE embedding model
- with IPEX-LLM optimizations on Intel CPUs and GPUs.
-
- To use, you should have the ``ipex-llm``
- and ``sentence_transformers`` package installed. Refer to
- `here `_
- for installation on Intel CPU.
-
- Example on Intel CPU:
- .. code-block:: python
-
- from langchain_community.embeddings import IpexLLMBgeEmbeddings
-
- embedding_model = IpexLLMBgeEmbeddings(
- model_name="BAAI/bge-large-en-v1.5",
- model_kwargs={},
- encode_kwargs={"normalize_embeddings": True},
- )
-
- Refer to
- `here `_
- for installation on Intel GPU.
-
- Example on Intel GPU:
- .. code-block:: python
-
- from langchain_community.embeddings import IpexLLMBgeEmbeddings
-
- embedding_model = IpexLLMBgeEmbeddings(
- model_name="BAAI/bge-large-en-v1.5",
- model_kwargs={"device": "xpu"},
- encode_kwargs={"normalize_embeddings": True},
- )
- """
-
- client: Any = None #: :meta private:
- model_name: str = DEFAULT_BGE_MODEL
- """Model name to use."""
- cache_folder: Optional[str] = None
- """Path to store models.
- Can be also set by SENTENCE_TRANSFORMERS_HOME environment variable."""
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Keyword arguments to pass to the model."""
- encode_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Keyword arguments to pass when calling the `encode` method of the model."""
- query_instruction: str = DEFAULT_QUERY_BGE_INSTRUCTION_EN
- """Instruction to use for embedding query."""
- embed_instruction: str = ""
- """Instruction to use for embedding document."""
-
- def __init__(self, **kwargs: Any):
- """Initialize the sentence_transformer."""
- super().__init__(**kwargs)
- try:
- import sentence_transformers
- from ipex_llm.transformers.convert import _optimize_post, _optimize_pre
-
- except ImportError as exc:
- base_url = (
- "https://python.langchain.com/v0.1/docs/integrations/text_embedding/"
- )
- raise ImportError(
- "Could not import ipex_llm or sentence_transformers. "
- f"Please refer to {base_url}/ipex_llm/ "
- "for install required packages on Intel CPU. "
- f"And refer to {base_url}/ipex_llm_gpu/ "
- "for install required packages on Intel GPU. "
- ) from exc
-
- # Set "cpu" as default device
- if "device" not in self.model_kwargs:
- self.model_kwargs["device"] = "cpu"
-
- if self.model_kwargs["device"] not in ["cpu", "xpu"]:
- raise ValueError(
- "IpexLLMBgeEmbeddings currently only supports device to be "
- f"'cpu' or 'xpu', but you have: {self.model_kwargs['device']}."
- )
-
- self.client = sentence_transformers.SentenceTransformer(
- self.model_name, cache_folder=self.cache_folder, **self.model_kwargs
- )
-
- # Add ipex-llm optimizations
- self.client = _optimize_pre(self.client)
- self.client = _optimize_post(self.client)
- if self.model_kwargs["device"] == "xpu":
- self.client = self.client.half().to("xpu")
-
- if "-zh" in self.model_name:
- self.query_instruction = DEFAULT_QUERY_BGE_INSTRUCTION_ZH
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Compute doc embeddings using a HuggingFace transformer model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- texts = [self.embed_instruction + t.replace("\n", " ") for t in texts]
- embeddings = self.client.encode(texts, **self.encode_kwargs)
- return embeddings.tolist()
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using a HuggingFace transformer model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- text = text.replace("\n", " ")
- embedding = self.client.encode(
- self.query_instruction + text, **self.encode_kwargs
- )
- return embedding.tolist()
diff --git a/libs/community/langchain_community/embeddings/itrex.py b/libs/community/langchain_community/embeddings/itrex.py
deleted file mode 100644
index 1f9a8e0731..0000000000
--- a/libs/community/langchain_community/embeddings/itrex.py
+++ /dev/null
@@ -1,214 +0,0 @@
-import importlib.util
-import os
-from typing import Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict
-
-
-class QuantizedBgeEmbeddings(BaseModel, Embeddings):
- """Leverage Itrex runtime to unlock the performance of compressed NLP models.
-
- Please ensure that you have installed intel-extension-for-transformers.
-
- Input:
- model_name: str = Model name.
- max_seq_len: int = The maximum sequence length for tokenization. (default 512)
- pooling_strategy: str =
- "mean" or "cls", pooling strategy for the final layer. (default "mean")
- query_instruction: Optional[str] =
- An instruction to add to the query before embedding. (default None)
- document_instruction: Optional[str] =
- An instruction to add to each document before embedding. (default None)
- padding: Optional[bool] =
- Whether to add padding during tokenization or not. (default True)
- model_kwargs: Optional[Dict] =
- Parameters to add to the model during initialization. (default {})
- encode_kwargs: Optional[Dict] =
- Parameters to add during the embedding forward pass. (default {})
- onnx_file_name: Optional[str] =
- File name of onnx optimized model which is exported by itrex.
- (default "int8-model.onnx")
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import QuantizedBgeEmbeddings
-
- model_name = "Intel/bge-small-en-v1.5-sts-int8-static-inc"
- encode_kwargs = {'normalize_embeddings': True}
- hf = QuantizedBgeEmbeddings(
- model_name,
- encode_kwargs=encode_kwargs,
- query_instruction="Represent this sentence for searching relevant passages: "
- )
- """ # noqa: E501
-
- def __init__(
- self,
- model_name: str,
- *,
- max_seq_len: int = 512,
- pooling_strategy: str = "mean", # "mean" or "cls"
- query_instruction: Optional[str] = None,
- document_instruction: Optional[str] = None,
- padding: bool = True,
- model_kwargs: Optional[Dict] = None,
- encode_kwargs: Optional[Dict] = None,
- onnx_file_name: Optional[str] = "int8-model.onnx",
- **kwargs: Any,
- ) -> None:
- super().__init__(**kwargs)
-
- # check sentence_transformers python package
- if importlib.util.find_spec("intel_extension_for_transformers") is None:
- raise ImportError(
- "Could not import intel_extension_for_transformers python package. "
- "Please install it with "
- "`pip install -U intel-extension-for-transformers`."
- )
-
- # check torch python package
- if importlib.util.find_spec("torch") is None:
- raise ImportError(
- "Could not import torch python package. "
- "Please install it with `pip install -U torch`."
- )
-
- # check onnx python package
- if importlib.util.find_spec("onnx") is None:
- raise ImportError(
- "Could not import onnx python package. "
- "Please install it with `pip install -U onnx`."
- )
-
- self.model_name_or_path = model_name
- self.max_seq_len = max_seq_len
- self.pooling = pooling_strategy
- self.padding = padding
- self.encode_kwargs = encode_kwargs or {}
- self.model_kwargs = model_kwargs or {}
-
- self.normalize = self.encode_kwargs.get("normalize_embeddings", False)
- self.batch_size = self.encode_kwargs.get("batch_size", 32)
-
- self.query_instruction = query_instruction
- self.document_instruction = document_instruction
- self.onnx_file_name = onnx_file_name
-
- self.load_model()
-
- def load_model(self) -> None:
- from huggingface_hub import hf_hub_download
- from intel_extension_for_transformers.transformers import AutoModel
- from transformers import AutoConfig, AutoTokenizer
-
- self.hidden_size = AutoConfig.from_pretrained(
- self.model_name_or_path
- ).hidden_size
- self.transformer_tokenizer = AutoTokenizer.from_pretrained(
- self.model_name_or_path,
- )
- onnx_model_path = os.path.join(self.model_name_or_path, self.onnx_file_name) # type: ignore[arg-type]
- if not os.path.exists(onnx_model_path):
- onnx_model_path = hf_hub_download(
- self.model_name_or_path, filename=self.onnx_file_name
- )
- self.transformer_model = AutoModel.from_pretrained(
- onnx_model_path, use_embedding_runtime=True
- )
-
- model_config = ConfigDict(
- extra="allow",
- protected_namespaces=(),
- )
-
- def _embed(self, inputs: Any) -> Any:
- import torch
-
- engine_input = [value for value in inputs.values()]
- outputs = self.transformer_model.generate(engine_input)
- if "last_hidden_state:0" in outputs:
- last_hidden_state = outputs["last_hidden_state:0"]
- else:
- last_hidden_state = [out for out in outputs.values()][0]
- last_hidden_state = torch.tensor(last_hidden_state).reshape(
- inputs["input_ids"].shape[0], inputs["input_ids"].shape[1], self.hidden_size
- )
- if self.pooling == "mean":
- emb = self._mean_pooling(last_hidden_state, inputs["attention_mask"])
- elif self.pooling == "cls":
- emb = self._cls_pooling(last_hidden_state)
- else:
- raise ValueError("pooling method no supported")
-
- if self.normalize:
- emb = torch.nn.functional.normalize(emb, p=2, dim=1)
- return emb
-
- @staticmethod
- def _cls_pooling(last_hidden_state: Any) -> Any:
- return last_hidden_state[:, 0]
-
- @staticmethod
- def _mean_pooling(last_hidden_state: Any, attention_mask: Any) -> Any:
- try:
- import torch
- except ImportError as e:
- raise ImportError(
- "Unable to import torch, please install with `pip install -U torch`."
- ) from e
- input_mask_expanded = (
- attention_mask.unsqueeze(-1).expand(last_hidden_state.size()).float()
- )
- sum_embeddings = torch.sum(last_hidden_state * input_mask_expanded, 1)
- sum_mask = torch.clamp(input_mask_expanded.sum(1), min=1e-9)
- return sum_embeddings / sum_mask
-
- def _embed_text(self, texts: List[str]) -> List[List[float]]:
- inputs = self.transformer_tokenizer(
- texts,
- max_length=self.max_seq_len,
- truncation=True,
- padding=self.padding,
- return_tensors="pt",
- )
- return self._embed(inputs).tolist()
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed a list of text documents using the Optimized Embedder model.
-
- Input:
- texts: List[str] = List of text documents to embed.
- Output:
- List[List[float]] = The embeddings of each text document.
- """
- try:
- import pandas as pd
- except ImportError as e:
- raise ImportError(
- "Unable to import pandas, please install with `pip install -U pandas`."
- ) from e
- docs = [
- self.document_instruction + d if self.document_instruction else d
- for d in texts
- ]
-
- # group into batches
- text_list_df = pd.DataFrame(docs, columns=["texts"]).reset_index()
-
- # assign each example with its batch
- text_list_df["batch_index"] = text_list_df["index"] // self.batch_size
-
- # create groups
- batches = list(text_list_df.groupby(["batch_index"])["texts"].apply(list))
-
- vectors = []
- for batch in batches:
- vectors += self._embed_text(batch)
- return vectors
-
- def embed_query(self, text: str) -> List[float]:
- if self.query_instruction:
- text = self.query_instruction + text
- return self._embed_text([text])[0]
diff --git a/libs/community/langchain_community/embeddings/javelin_ai_gateway.py b/libs/community/langchain_community/embeddings/javelin_ai_gateway.py
deleted file mode 100644
index 205e58e1c3..0000000000
--- a/libs/community/langchain_community/embeddings/javelin_ai_gateway.py
+++ /dev/null
@@ -1,109 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Iterator, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel
-
-
-def _chunk(texts: List[str], size: int) -> Iterator[List[str]]:
- for i in range(0, len(texts), size):
- yield texts[i : i + size]
-
-
-class JavelinAIGatewayEmbeddings(Embeddings, BaseModel):
- """Javelin AI Gateway embeddings.
-
- To use, you should have the ``javelin_sdk`` python package installed.
- For more information, see https://docs.getjavelin.io
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import JavelinAIGatewayEmbeddings
-
- embeddings = JavelinAIGatewayEmbeddings(
- gateway_uri="",
- route=""
- )
- """
-
- client: Any
- """javelin client."""
-
- route: str
- """The route to use for the Javelin AI Gateway API."""
-
- gateway_uri: Optional[str] = None
- """The URI for the Javelin AI Gateway API."""
-
- javelin_api_key: Optional[str] = None
- """The API key for the Javelin AI Gateway API."""
-
- def __init__(self, **kwargs: Any):
- try:
- from javelin_sdk import (
- JavelinClient,
- UnauthorizedError,
- )
- except ImportError:
- raise ImportError(
- "Could not import javelin_sdk python package. "
- "Please install it with `pip install javelin_sdk`."
- )
-
- super().__init__(**kwargs)
- if self.gateway_uri:
- try:
- self.client = JavelinClient(
- base_url=self.gateway_uri, api_key=self.javelin_api_key
- )
- except UnauthorizedError as e:
- raise ValueError("Javelin: Incorrect API Key.") from e
-
- def _query(self, texts: List[str]) -> List[List[float]]:
- embeddings = []
- for txt in _chunk(texts, 20):
- try:
- resp = self.client.query_route(self.route, query_body={"input": txt})
- resp_dict = resp.dict()
-
- embeddings_chunk = resp_dict.get("llm_response", {}).get("data", [])
- for item in embeddings_chunk:
- if "embedding" in item:
- embeddings.append(item["embedding"])
- except ValueError as e:
- print("Failed to query route: " + str(e)) # noqa: T201
-
- return embeddings
-
- async def _aquery(self, texts: List[str]) -> List[List[float]]:
- embeddings = []
- for txt in _chunk(texts, 20):
- try:
- resp = await self.client.aquery_route(
- self.route, query_body={"input": txt}
- )
- resp_dict = resp.dict()
-
- embeddings_chunk = resp_dict.get("llm_response", {}).get("data", [])
- for item in embeddings_chunk:
- if "embedding" in item:
- embeddings.append(item["embedding"])
- except ValueError as e:
- print("Failed to query route: " + str(e)) # noqa: T201
-
- return embeddings
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- return self._query(texts)
-
- def embed_query(self, text: str) -> List[float]:
- return self._query([text])[0]
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- return await self._aquery(texts)
-
- async def aembed_query(self, text: str) -> List[float]:
- result = await self._aquery([text])
- return result[0]
diff --git a/libs/community/langchain_community/embeddings/jina.py b/libs/community/langchain_community/embeddings/jina.py
deleted file mode 100644
index ad9ea9fd92..0000000000
--- a/libs/community/langchain_community/embeddings/jina.py
+++ /dev/null
@@ -1,124 +0,0 @@
-import base64
-from os.path import exists
-from typing import Any, Dict, List, Optional
-from urllib.parse import urlparse
-
-import requests
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env
-from pydantic import BaseModel, ConfigDict, SecretStr, model_validator
-
-JINA_API_URL: str = "https://api.jina.ai/v1/embeddings"
-
-
-def is_local(url: str) -> bool:
- """Check if a URL is a local file.
-
- Args:
- url (str): The URL to check.
-
- Returns:
- bool: True if the URL is a local file, False otherwise.
- """
- url_parsed = urlparse(url)
- if url_parsed.scheme in ("file", ""): # Possibly a local file
- return exists(url_parsed.path)
- return False
-
-
-def get_bytes_str(file_path: str) -> str:
- """Get the bytes string of a file.
-
- Args:
- file_path (str): The path to the file.
-
- Returns:
- str: The bytes string of the file.
- """
- with open(file_path, "rb") as image_file:
- return base64.b64encode(image_file.read()).decode("utf-8")
-
-
-class JinaEmbeddings(BaseModel, Embeddings):
- """Jina embedding models."""
-
- session: Any #: :meta private:
- model_name: str = "jina-embeddings-v2-base-en"
- jina_api_key: Optional[SecretStr] = None
-
- model_config = ConfigDict(protected_namespaces=())
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that auth token exists in environment."""
- try:
- jina_api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "jina_api_key", "JINA_API_KEY")
- )
- except ValueError as original_exc:
- try:
- jina_api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "jina_auth_token", "JINA_AUTH_TOKEN")
- )
- except ValueError:
- raise original_exc
- session = requests.Session()
- session.headers.update(
- {
- "Authorization": f"Bearer {jina_api_key.get_secret_value()}",
- "Accept-Encoding": "identity",
- "Content-type": "application/json",
- }
- )
- values["session"] = session
- return values
-
- def _embed(self, input: Any) -> List[List[float]]:
- # Call Jina AI Embedding API
- resp = self.session.post(
- JINA_API_URL, json={"input": input, "model": self.model_name}
- ).json()
- if "data" not in resp:
- raise RuntimeError(resp["detail"])
-
- embeddings = resp["data"]
-
- # Sort resulting embeddings by index
- sorted_embeddings = sorted(embeddings, key=lambda e: e["index"])
-
- # Return just the embeddings
- return [result["embedding"] for result in sorted_embeddings]
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Call out to Jina's embedding endpoint.
- Args:
- texts: The list of texts to embed.
- Returns:
- List of embeddings, one for each text.
- """
- return self._embed(texts)
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to Jina's embedding endpoint.
- Args:
- text: The text to embed.
- Returns:
- Embeddings for the text.
- """
- return self._embed([text])[0]
-
- def embed_images(self, uris: List[str]) -> List[List[float]]:
- """Call out to Jina's image embedding endpoint.
- Args:
- uris: The list of uris to embed.
- Returns:
- List of embeddings, one for each text.
- """
- input = []
- for uri in uris:
- if is_local(uri):
- input.append({"bytes": get_bytes_str(uri)})
- else:
- input.append({"url": uri})
- return self._embed(input)
diff --git a/libs/community/langchain_community/embeddings/johnsnowlabs.py b/libs/community/langchain_community/embeddings/johnsnowlabs.py
deleted file mode 100644
index 4223114aa0..0000000000
--- a/libs/community/langchain_community/embeddings/johnsnowlabs.py
+++ /dev/null
@@ -1,91 +0,0 @@
-import os
-import sys
-from typing import Any, List
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict
-
-
-class JohnSnowLabsEmbeddings(BaseModel, Embeddings):
- """JohnSnowLabs embedding models
-
- To use, you should have the ``johnsnowlabs`` python package installed.
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings.johnsnowlabs import JohnSnowLabsEmbeddings
-
- embedding = JohnSnowLabsEmbeddings(model='embed_sentence.bert')
- output = embedding.embed_query("foo bar")
- """ # noqa: E501
-
- model: Any = "embed_sentence.bert"
-
- def __init__(
- self,
- model: Any = "embed_sentence.bert",
- hardware_target: str = "cpu",
- **kwargs: Any,
- ):
- """Initialize the johnsnowlabs model."""
- super().__init__(**kwargs)
- # 1) Check imports
- try:
- from johnsnowlabs import nlp
- from nlu.pipe.pipeline import NLUPipeline
- except ImportError as exc:
- raise ImportError(
- "Could not import johnsnowlabs python package. "
- "Please install it with `pip install johnsnowlabs`."
- ) from exc
-
- # 2) Start a Spark Session
- try:
- os.environ["PYSPARK_PYTHON"] = sys.executable
- os.environ["PYSPARK_DRIVER_PYTHON"] = sys.executable
- nlp.start(hardware_target=hardware_target)
- except Exception as exc:
- raise Exception("Failure starting Spark Session") from exc
-
- # 3) Load the model
- try:
- if isinstance(model, str):
- self.model = nlp.load(model)
- elif isinstance(model, NLUPipeline):
- self.model = model
- else:
- self.model = nlp.to_nlu_pipe(model)
- except Exception as exc:
- raise Exception("Failure loading model") from exc
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Compute doc embeddings using a JohnSnowLabs transformer model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
-
- df = self.model.predict(texts, output_level="document")
- emb_col = None
- for c in df.columns:
- if "embedding" in c:
- emb_col = c
- return [vec.tolist() for vec in df[emb_col].tolist()]
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using a JohnSnowLabs transformer model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self.embed_documents([text])[0]
diff --git a/libs/community/langchain_community/embeddings/laser.py b/libs/community/langchain_community/embeddings/laser.py
deleted file mode 100644
index 088ffbb4d1..0000000000
--- a/libs/community/langchain_community/embeddings/laser.py
+++ /dev/null
@@ -1,89 +0,0 @@
-from typing import Any, Dict, List, Optional, cast
-
-import numpy as np
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import pre_init
-from pydantic import BaseModel, ConfigDict
-
-LASER_MULTILINGUAL_MODEL: str = "laser2"
-
-
-class LaserEmbeddings(BaseModel, Embeddings):
- """LASER Language-Agnostic SEntence Representations.
- LASER is a Python library developed by the Meta AI Research team
- and used for creating multilingual sentence embeddings for over 147 languages
- as of 2/25/2024
- See more documentation at:
- * https://github.com/facebookresearch/LASER/
- * https://github.com/facebookresearch/LASER/tree/main/laser_encoders
- * https://arxiv.org/abs/2205.12654
-
- To use this class, you must install the `laser_encoders` Python package.
-
- `pip install laser_encoders`
- Example:
- from laser_encoders import LaserEncoderPipeline
- encoder = LaserEncoderPipeline(lang="eng_Latn")
- embeddings = encoder.encode_sentences(["Hello", "World"])
- """
-
- lang: Optional[str] = None
- """The language or language code you'd like to use
- If empty, this implementation will default
- to using a multilingual earlier LASER encoder model (called laser2)
- Find the list of supported languages at
- https://github.com/facebookresearch/flores/blob/main/flores200/README.md#languages-in-flores-200
- """
-
- _encoder_pipeline: Any = None # : :meta private:
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that laser_encoders has been installed."""
- try:
- from laser_encoders import LaserEncoderPipeline
-
- lang = values.get("lang")
- if lang:
- encoder_pipeline = LaserEncoderPipeline(lang=lang)
- else:
- encoder_pipeline = LaserEncoderPipeline(laser=LASER_MULTILINGUAL_MODEL)
- values["_encoder_pipeline"] = encoder_pipeline
-
- except ImportError as e:
- raise ImportError(
- "Could not import 'laser_encoders' Python package. "
- "Please install it with `pip install laser_encoders`."
- ) from e
- return values
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Generate embeddings for documents using LASER.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- embeddings: np.ndarray
- embeddings = self._encoder_pipeline.encode_sentences(texts)
-
- return cast(List[List[float]], embeddings.tolist())
-
- def embed_query(self, text: str) -> List[float]:
- """Generate single query text embeddings using LASER.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- query_embeddings: np.ndarray
- query_embeddings = self._encoder_pipeline.encode_sentences([text])
- return cast(List[List[float]], query_embeddings.tolist())[0]
diff --git a/libs/community/langchain_community/embeddings/llamacpp.py b/libs/community/langchain_community/embeddings/llamacpp.py
deleted file mode 100644
index e4ebe33b33..0000000000
--- a/libs/community/langchain_community/embeddings/llamacpp.py
+++ /dev/null
@@ -1,145 +0,0 @@
-from typing import Any, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict, Field, model_validator
-from typing_extensions import Self
-
-
-class LlamaCppEmbeddings(BaseModel, Embeddings):
- """llama.cpp embedding models.
-
- To use, you should have the llama-cpp-python library installed, and provide the
- path to the Llama model as a named parameter to the constructor.
- Check out: https://github.com/abetlen/llama-cpp-python
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import LlamaCppEmbeddings
- llama = LlamaCppEmbeddings(model_path="/path/to/model.bin")
- """
-
- client: Any = None #: :meta private:
- model_path: str = Field(default="")
-
- n_ctx: int = Field(512, alias="n_ctx")
- """Token context window."""
-
- n_parts: int = Field(-1, alias="n_parts")
- """Number of parts to split the model into.
- If -1, the number of parts is automatically determined."""
-
- seed: int = Field(-1, alias="seed")
- """Seed. If -1, a random seed is used."""
-
- f16_kv: bool = Field(False, alias="f16_kv")
- """Use half-precision for key/value cache."""
-
- logits_all: bool = Field(False, alias="logits_all")
- """Return logits for all tokens, not just the last token."""
-
- vocab_only: bool = Field(False, alias="vocab_only")
- """Only load the vocabulary, no weights."""
-
- use_mlock: bool = Field(False, alias="use_mlock")
- """Force system to keep model in RAM."""
-
- n_threads: Optional[int] = Field(None, alias="n_threads")
- """Number of threads to use. If None, the number
- of threads is automatically determined."""
-
- n_batch: Optional[int] = Field(512, alias="n_batch")
- """Number of tokens to process in parallel.
- Should be a number between 1 and n_ctx."""
-
- n_gpu_layers: Optional[int] = Field(None, alias="n_gpu_layers")
- """Number of layers to be loaded into gpu memory. Default None."""
-
- verbose: bool = Field(True, alias="verbose")
- """Print verbose output to stderr."""
-
- device: Optional[str] = Field(None, alias="device")
- """Device type to use and pass to the model"""
-
- model_config = ConfigDict(
- extra="forbid",
- protected_namespaces=(),
- )
-
- @model_validator(mode="after")
- def validate_environment(self) -> Self:
- """Validate that llama-cpp-python library is installed."""
- model_path = self.model_path
- model_param_names = [
- "n_ctx",
- "n_parts",
- "seed",
- "f16_kv",
- "logits_all",
- "vocab_only",
- "use_mlock",
- "n_threads",
- "n_batch",
- "verbose",
- "device",
- ]
- model_params = {k: getattr(self, k) for k in model_param_names}
- # For backwards compatibility, only include if non-null.
- if self.n_gpu_layers is not None:
- model_params["n_gpu_layers"] = self.n_gpu_layers
-
- if not self.client:
- try:
- from llama_cpp import Llama
-
- self.client = Llama(model_path, embedding=True, **model_params)
- except ImportError:
- raise ImportError(
- "Could not import llama-cpp-python library. "
- "Please install the llama-cpp-python library to "
- "use this embedding model: pip install llama-cpp-python"
- )
- except Exception as e:
- raise ValueError(
- f"Could not load Llama model from path: {model_path}. "
- f"Received error {e}"
- )
-
- return self
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed a list of documents using the Llama model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- embeddings = self.client.create_embedding(texts)
- final_embeddings = []
- for e in embeddings["data"]:
- try:
- if isinstance(e["embedding"][0], list):
- for data in e["embedding"]:
- final_embeddings.append(list(map(float, data)))
- else:
- final_embeddings.append(list(map(float, e["embedding"])))
- except (IndexError, TypeError):
- final_embeddings.append(list(map(float, e["embedding"])))
- return final_embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using the Llama model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- embedding = self.client.embed(text)
- if embedding and isinstance(embedding, list) and isinstance(embedding[0], list):
- return list(map(float, embedding[0]))
- else:
- return list(map(float, embedding))
diff --git a/libs/community/langchain_community/embeddings/llamafile.py b/libs/community/langchain_community/embeddings/llamafile.py
deleted file mode 100644
index 247b1a923a..0000000000
--- a/libs/community/langchain_community/embeddings/llamafile.py
+++ /dev/null
@@ -1,119 +0,0 @@
-import logging
-from typing import List, Optional
-
-import requests
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel
-
-logger = logging.getLogger(__name__)
-
-
-class LlamafileEmbeddings(BaseModel, Embeddings):
- """Llamafile lets you distribute and run large language models with a
- single file.
-
- To get started, see: https://github.com/Mozilla-Ocho/llamafile
-
- To use this class, you will need to first:
-
- 1. Download a llamafile.
- 2. Make the downloaded file executable: `chmod +x path/to/model.llamafile`
- 3. Start the llamafile in server mode with embeddings enabled:
-
- `./path/to/model.llamafile --server --nobrowser --embedding`
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import LlamafileEmbeddings
- embedder = LlamafileEmbeddings()
- doc_embeddings = embedder.embed_documents(
- [
- "Alpha is the first letter of the Greek alphabet",
- "Beta is the second letter of the Greek alphabet",
- ]
- )
- query_embedding = embedder.embed_query(
- "What is the second letter of the Greek alphabet"
- )
-
- """
-
- base_url: str = "http://localhost:8080"
- """Base url where the llamafile server is listening."""
-
- request_timeout: Optional[int] = None
- """Timeout for server requests"""
-
- def _embed(self, text: str) -> List[float]:
- try:
- response = requests.post(
- url=f"{self.base_url}/embedding",
- headers={
- "Content-Type": "application/json",
- },
- json={
- "content": text,
- },
- timeout=self.request_timeout,
- )
- except requests.exceptions.ConnectionError:
- raise requests.exceptions.ConnectionError(
- f"Could not connect to Llamafile server. Please make sure "
- f"that a server is running at {self.base_url}."
- )
-
- # Raise exception if we got a bad (non-200) response status code
- response.raise_for_status()
-
- contents = response.json()
- if "embedding" not in contents:
- raise KeyError(
- "Unexpected output from /embedding endpoint, output dict "
- "missing 'embedding' key."
- )
-
- embedding = contents["embedding"]
-
- # Sanity check the embedding vector:
- # Prior to llamafile v0.6.2, if the server was not started with the
- # `--embedding` option, the embedding endpoint would always return a
- # 0-vector. See issue:
- # https://github.com/Mozilla-Ocho/llamafile/issues/243
- # So here we raise an exception if the vector sums to exactly 0.
- if sum(embedding) == 0.0:
- raise ValueError(
- "Embedding sums to 0, did you start the llamafile server with "
- "the `--embedding` option enabled?"
- )
-
- return embedding
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed documents using a llamafile server running at `self.base_url`.
- llamafile server should be started in a separate process before invoking
- this method.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- doc_embeddings = []
- for text in texts:
- doc_embeddings.append(self._embed(text))
- return doc_embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using a llamafile server running at `self.base_url`.
- llamafile server should be started in a separate process before invoking
- this method.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self._embed(text)
diff --git a/libs/community/langchain_community/embeddings/llm_rails.py b/libs/community/langchain_community/embeddings/llm_rails.py
deleted file mode 100644
index 92bb8c6a10..0000000000
--- a/libs/community/langchain_community/embeddings/llm_rails.py
+++ /dev/null
@@ -1,74 +0,0 @@
-"""This file is for LLMRails Embedding"""
-
-from typing import Dict, List, Optional
-
-import requests
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict, SecretStr
-
-
-class LLMRailsEmbeddings(BaseModel, Embeddings):
- """LLMRails embedding models.
-
- To use, you should have the environment
- variable ``LLM_RAILS_API_KEY`` set with your API key or pass it
- as a named parameter to the constructor.
-
- Model can be one of ["embedding-english-v1","embedding-multi-v1"]
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import LLMRailsEmbeddings
- cohere = LLMRailsEmbeddings(
- model="embedding-english-v1", api_key="my-api-key"
- )
- """
-
- model: str = "embedding-english-v1"
- """Model name to use."""
-
- api_key: Optional[SecretStr] = None
- """LLMRails API key."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key exists in environment."""
- api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "api_key", "LLM_RAILS_API_KEY")
- )
- values["api_key"] = api_key
- return values
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Call out to Cohere's embedding endpoint.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- response = requests.post(
- "https://api.llmrails.com/v1/embeddings",
- headers={"X-API-KEY": self.api_key.get_secret_value()}, # type: ignore[union-attr]
- json={"input": texts, "model": self.model},
- timeout=60,
- )
- return [item["embedding"] for item in response.json()["data"]]
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to Cohere's embedding endpoint.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self.embed_documents([text])[0]
diff --git a/libs/community/langchain_community/embeddings/localai.py b/libs/community/langchain_community/embeddings/localai.py
deleted file mode 100644
index 8c0457b395..0000000000
--- a/libs/community/langchain_community/embeddings/localai.py
+++ /dev/null
@@ -1,347 +0,0 @@
-from __future__ import annotations
-
-import logging
-import warnings
-from typing import (
- Any,
- Callable,
- Dict,
- List,
- Literal,
- Optional,
- Sequence,
- Set,
- Tuple,
- Union,
-)
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import (
- get_from_dict_or_env,
- get_pydantic_field_names,
- pre_init,
-)
-from pydantic import BaseModel, ConfigDict, Field, model_validator
-from tenacity import (
- AsyncRetrying,
- before_sleep_log,
- retry,
- retry_if_exception_type,
- stop_after_attempt,
- wait_exponential,
-)
-
-logger = logging.getLogger(__name__)
-
-
-def _create_retry_decorator(embeddings: LocalAIEmbeddings) -> Callable[[Any], Any]:
- import openai
-
- min_seconds = 4
- max_seconds = 10
- # Wait 2^x * 1 second between each retry starting with
- # 4 seconds, then up to 10 seconds, then 10 seconds afterwards
- return retry(
- reraise=True,
- stop=stop_after_attempt(embeddings.max_retries),
- wait=wait_exponential(multiplier=1, min=min_seconds, max=max_seconds),
- retry=(
- retry_if_exception_type(openai.error.Timeout)
- | retry_if_exception_type(openai.error.APIError)
- | retry_if_exception_type(openai.error.APIConnectionError)
- | retry_if_exception_type(openai.error.RateLimitError)
- | retry_if_exception_type(openai.error.ServiceUnavailableError)
- ),
- before_sleep=before_sleep_log(logger, logging.WARNING),
- )
-
-
-def _async_retry_decorator(embeddings: LocalAIEmbeddings) -> Any:
- import openai
-
- min_seconds = 4
- max_seconds = 10
- # Wait 2^x * 1 second between each retry starting with
- # 4 seconds, then up to 10 seconds, then 10 seconds afterwards
- async_retrying = AsyncRetrying(
- reraise=True,
- stop=stop_after_attempt(embeddings.max_retries),
- wait=wait_exponential(multiplier=1, min=min_seconds, max=max_seconds),
- retry=(
- retry_if_exception_type(openai.error.Timeout)
- | retry_if_exception_type(openai.error.APIError)
- | retry_if_exception_type(openai.error.APIConnectionError)
- | retry_if_exception_type(openai.error.RateLimitError)
- | retry_if_exception_type(openai.error.ServiceUnavailableError)
- ),
- before_sleep=before_sleep_log(logger, logging.WARNING),
- )
-
- def wrap(func: Callable) -> Callable:
- async def wrapped_f(*args: Any, **kwargs: Any) -> Callable:
- async for _ in async_retrying:
- return await func(*args, **kwargs)
- raise AssertionError("this is unreachable")
-
- return wrapped_f
-
- return wrap
-
-
-# https://stackoverflow.com/questions/76469415/getting-embeddings-of-length-1-from-langchain-openaiembeddings
-def _check_response(response: dict) -> dict:
- if any(len(d["embedding"]) == 1 for d in response["data"]):
- import openai
-
- raise openai.error.APIError("LocalAI API returned an empty embedding")
- return response
-
-
-def embed_with_retry(embeddings: LocalAIEmbeddings, **kwargs: Any) -> Any:
- """Use tenacity to retry the embedding call."""
- retry_decorator = _create_retry_decorator(embeddings)
-
- @retry_decorator
- def _embed_with_retry(**kwargs: Any) -> Any:
- response = embeddings.client.create(**kwargs)
- return _check_response(response)
-
- return _embed_with_retry(**kwargs)
-
-
-async def async_embed_with_retry(embeddings: LocalAIEmbeddings, **kwargs: Any) -> Any:
- """Use tenacity to retry the embedding call."""
-
- @_async_retry_decorator(embeddings)
- async def _async_embed_with_retry(**kwargs: Any) -> Any:
- response = await embeddings.client.acreate(**kwargs)
- return _check_response(response)
-
- return await _async_embed_with_retry(**kwargs)
-
-
-class LocalAIEmbeddings(BaseModel, Embeddings):
- """LocalAI embedding models.
-
- Since LocalAI and OpenAI have 1:1 compatibility between APIs, this class
- uses the ``openai`` Python package's ``openai.Embedding`` as its client.
- Thus, you should have the ``openai`` python package installed, and defeat
- the environment variable ``OPENAI_API_KEY`` by setting to a random string.
- You also need to specify ``OPENAI_API_BASE`` to point to your LocalAI
- service endpoint.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import LocalAIEmbeddings
- openai = LocalAIEmbeddings(
- openai_api_key="random-string",
- openai_api_base="http://localhost:8080"
- )
-
- """
-
- client: Any = None #: :meta private:
- model: str = "text-embedding-ada-002"
- deployment: str = model
- openai_api_version: Optional[str] = None
- openai_api_base: Optional[str] = None
- # to support explicit proxy for LocalAI
- openai_proxy: Optional[str] = None
- embedding_ctx_length: int = 8191
- """The maximum number of tokens to embed at once."""
- openai_api_key: Optional[str] = None
- openai_organization: Optional[str] = None
- allowed_special: Union[Literal["all"], Set[str]] = set()
- disallowed_special: Union[Literal["all"], Set[str], Sequence[str]] = "all"
- chunk_size: int = 1000
- """Maximum number of texts to embed in each batch"""
- max_retries: int = 6
- """Maximum number of retries to make when generating."""
- request_timeout: Optional[Union[float, Tuple[float, float]]] = None
- """Timeout in seconds for the LocalAI request."""
- headers: Any = None
- show_progress_bar: bool = False
- """Whether to show a progress bar when embedding."""
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `create` call not explicitly specified."""
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = get_pydantic_field_names(cls)
- extra = values.get("model_kwargs", {})
- for field_name in list(values):
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- if field_name not in all_required_field_names:
- warnings.warn(
- f"""WARNING! {field_name} is not default parameter.
- {field_name} was transferred to model_kwargs.
- Please confirm that {field_name} is what you intended."""
- )
- extra[field_name] = values.pop(field_name)
-
- invalid_model_kwargs = all_required_field_names.intersection(extra.keys())
- if invalid_model_kwargs:
- raise ValueError(
- f"Parameters {invalid_model_kwargs} should be specified explicitly. "
- f"Instead they were passed in as part of `model_kwargs` parameter."
- )
-
- values["model_kwargs"] = extra
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- values["openai_api_key"] = get_from_dict_or_env(
- values, "openai_api_key", "OPENAI_API_KEY"
- )
- values["openai_api_base"] = get_from_dict_or_env(
- values,
- "openai_api_base",
- "OPENAI_API_BASE",
- default="",
- )
- values["openai_proxy"] = get_from_dict_or_env(
- values,
- "openai_proxy",
- "OPENAI_PROXY",
- default="",
- )
-
- default_api_version = ""
- values["openai_api_version"] = get_from_dict_or_env(
- values,
- "openai_api_version",
- "OPENAI_API_VERSION",
- default=default_api_version,
- )
- values["openai_organization"] = get_from_dict_or_env(
- values,
- "openai_organization",
- "OPENAI_ORGANIZATION",
- default="",
- )
- try:
- import openai
-
- values["client"] = openai.Embedding
- except ImportError:
- raise ImportError(
- "Could not import openai python package. "
- "Please install it with `pip install openai`."
- )
- return values
-
- @property
- def _invocation_params(self) -> Dict:
- openai_args = {
- "model": self.model,
- "request_timeout": self.request_timeout,
- "headers": self.headers,
- "api_key": self.openai_api_key,
- "organization": self.openai_organization,
- "api_base": self.openai_api_base,
- "api_version": self.openai_api_version,
- **self.model_kwargs,
- }
- if self.openai_proxy:
- import openai
-
- openai.proxy = {
- "http": self.openai_proxy,
- "https": self.openai_proxy,
- }
- return openai_args
-
- def _embedding_func(self, text: str, *, engine: str) -> List[float]:
- """Call out to LocalAI's embedding endpoint."""
- # handle large input text
- if self.model.endswith("001"):
- # See: https://github.com/openai/openai-python/issues/418#issuecomment-1525939500
- # replace newlines, which can negatively affect performance.
- text = text.replace("\n", " ")
- return embed_with_retry(
- self,
- input=[text],
- **self._invocation_params,
- )["data"][0]["embedding"]
-
- async def _aembedding_func(self, text: str, *, engine: str) -> List[float]:
- """Call out to LocalAI's embedding endpoint."""
- # handle large input text
- if self.model.endswith("001"):
- # See: https://github.com/openai/openai-python/issues/418#issuecomment-1525939500
- # replace newlines, which can negatively affect performance.
- text = text.replace("\n", " ")
- return (
- await async_embed_with_retry(
- self,
- input=[text],
- **self._invocation_params,
- )
- )["data"][0]["embedding"]
-
- def embed_documents(
- self, texts: List[str], chunk_size: Optional[int] = 0
- ) -> List[List[float]]:
- """Call out to LocalAI's embedding endpoint for embedding search docs.
-
- Args:
- texts: The list of texts to embed.
- chunk_size: The chunk size of embeddings. If None, will use the chunk size
- specified by the class.
-
- Returns:
- List of embeddings, one for each text.
- """
- # call _embedding_func for each text
- return [self._embedding_func(text, engine=self.deployment) for text in texts]
-
- async def aembed_documents(
- self, texts: List[str], chunk_size: Optional[int] = 0
- ) -> List[List[float]]:
- """Call out to LocalAI's embedding endpoint async for embedding search docs.
-
- Args:
- texts: The list of texts to embed.
- chunk_size: The chunk size of embeddings. If None, will use the chunk size
- specified by the class.
-
- Returns:
- List of embeddings, one for each text.
- """
- embeddings = []
- for text in texts:
- response = await self._aembedding_func(text, engine=self.deployment)
- embeddings.append(response)
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to LocalAI's embedding endpoint for embedding query text.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embedding for the text.
- """
- embedding = self._embedding_func(text, engine=self.deployment)
- return embedding
-
- async def aembed_query(self, text: str) -> List[float]:
- """Call out to LocalAI's embedding endpoint async for embedding query text.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embedding for the text.
- """
- embedding = await self._aembedding_func(text, engine=self.deployment)
- return embedding
diff --git a/libs/community/langchain_community/embeddings/minimax.py b/libs/community/langchain_community/embeddings/minimax.py
deleted file mode 100644
index 1426278615..0000000000
--- a/libs/community/langchain_community/embeddings/minimax.py
+++ /dev/null
@@ -1,201 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Callable, Dict, List, Optional
-
-import requests
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict, Field, SecretStr
-from tenacity import (
- before_sleep_log,
- retry,
- stop_after_attempt,
- wait_exponential,
-)
-
-logger = logging.getLogger(__name__)
-
-
-def _create_retry_decorator() -> Callable[[Any], Any]:
- """Returns a tenacity retry decorator."""
-
- multiplier = 1
- min_seconds = 1
- max_seconds = 4
- max_retries = 6
-
- return retry(
- reraise=True,
- stop=stop_after_attempt(max_retries),
- wait=wait_exponential(multiplier=multiplier, min=min_seconds, max=max_seconds),
- before_sleep=before_sleep_log(logger, logging.WARNING),
- )
-
-
-def embed_with_retry(embeddings: MiniMaxEmbeddings, *args: Any, **kwargs: Any) -> Any:
- """Use tenacity to retry the completion call."""
- retry_decorator = _create_retry_decorator()
-
- @retry_decorator
- def _embed_with_retry(*args: Any, **kwargs: Any) -> Any:
- return embeddings.embed(*args, **kwargs)
-
- return _embed_with_retry(*args, **kwargs)
-
-
-class MiniMaxEmbeddings(BaseModel, Embeddings):
- """MiniMax embedding model integration.
-
- Setup:
- To use, you should have the environment variable ``MINIMAX_GROUP_ID`` and
- ``MINIMAX_API_KEY`` set with your API token.
-
- .. code-block:: bash
-
- export MINIMAX_API_KEY="your-api-key"
- export MINIMAX_GROUP_ID="your-group-id"
-
- Key init args — completion params:
- model: Optional[str]
- Name of ZhipuAI model to use.
- api_key: Optional[str]
- Automatically inferred from env var `MINIMAX_GROUP_ID` if not provided.
- group_id: Optional[str]
- Automatically inferred from env var `MINIMAX_GROUP_ID` if not provided.
-
- See full list of supported init args and their descriptions in the params section.
-
- Instantiate:
-
- .. code-block:: python
-
- from langchain_community.embeddings import MiniMaxEmbeddings
-
- embed = MiniMaxEmbeddings(
- model="embo-01",
- # api_key="...",
- # group_id="...",
- # other
- )
-
- Embed single text:
- .. code-block:: python
-
- input_text = "The meaning of life is 42"
- embed.embed_query(input_text)
-
- .. code-block:: python
-
- [0.03016241, 0.03617699, 0.0017198119, -0.002061239, -0.00029994643, -0.0061320597, -0.0043635326, ...]
-
- Embed multiple text:
- .. code-block:: python
-
- input_texts = ["This is a test query1.", "This is a test query2."]
- embed.embed_documents(input_texts)
-
- .. code-block:: python
-
- [
- [-0.0021588828, -0.007608119, 0.029349545, -0.0038194496, 0.008031177, -0.004529633, -0.020150753, ...],
- [ -0.00023150232, -0.011122423, 0.016930554, 0.0083089275, 0.012633711, 0.019683322, -0.005971041, ...]
- ]
- """ # noqa: E501
-
- endpoint_url: str = "https://api.minimax.chat/v1/embeddings"
- """Endpoint URL to use."""
- model: str = "embo-01"
- """Embeddings model name to use."""
- embed_type_db: str = "db"
- """For embed_documents"""
- embed_type_query: str = "query"
- """For embed_query"""
-
- minimax_group_id: Optional[str] = Field(default=None, alias="group_id")
- """Group ID for MiniMax API."""
- minimax_api_key: Optional[SecretStr] = Field(default=None, alias="api_key")
- """API Key for MiniMax API."""
-
- model_config = ConfigDict(
- populate_by_name=True,
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that group id and api key exists in environment."""
- minimax_group_id = get_from_dict_or_env(
- values, ["minimax_group_id", "group_id"], "MINIMAX_GROUP_ID"
- )
- minimax_api_key = convert_to_secret_str(
- get_from_dict_or_env(
- values, ["minimax_api_key", "api_key"], "MINIMAX_API_KEY"
- )
- )
- values["minimax_group_id"] = minimax_group_id
- values["minimax_api_key"] = minimax_api_key
- return values
-
- def embed(
- self,
- texts: List[str],
- embed_type: str,
- ) -> List[List[float]]:
- payload = {
- "model": self.model,
- "type": embed_type,
- "texts": texts,
- }
-
- # HTTP headers for authorization
- headers = {
- "Authorization": f"Bearer {self.minimax_api_key.get_secret_value()}", # type: ignore[union-attr]
- "Content-Type": "application/json",
- }
-
- params = {
- "GroupId": self.minimax_group_id,
- }
-
- # send request
- response = requests.post(
- self.endpoint_url, params=params, headers=headers, json=payload
- )
- parsed_response = response.json()
-
- # check for errors
- if parsed_response["base_resp"]["status_code"] != 0:
- raise ValueError(
- f"MiniMax API returned an error: {parsed_response['base_resp']}"
- )
-
- embeddings = parsed_response["vectors"]
-
- return embeddings
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed documents using a MiniMax embedding endpoint.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- embeddings = embed_with_retry(self, texts=texts, embed_type=self.embed_type_db)
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using a MiniMax embedding endpoint.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- embeddings = embed_with_retry(
- self, texts=[text], embed_type=self.embed_type_query
- )
- return embeddings[0]
diff --git a/libs/community/langchain_community/embeddings/mlflow.py b/libs/community/langchain_community/embeddings/mlflow.py
deleted file mode 100644
index 09ceb3a229..0000000000
--- a/libs/community/langchain_community/embeddings/mlflow.py
+++ /dev/null
@@ -1,91 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Dict, Iterator, List
-from urllib.parse import urlparse
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, PrivateAttr
-
-
-def _chunk(texts: List[str], size: int) -> Iterator[List[str]]:
- for i in range(0, len(texts), size):
- yield texts[i : i + size]
-
-
-class MlflowEmbeddings(Embeddings, BaseModel):
- """Embedding LLMs in MLflow.
-
- To use, you should have the `mlflow[genai]` python package installed.
- For more information, see https://mlflow.org/docs/latest/llms/deployments.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import MlflowEmbeddings
-
- embeddings = MlflowEmbeddings(
- target_uri="http://localhost:5000",
- endpoint="embeddings",
- )
- """
-
- endpoint: str
- """The endpoint to use."""
- target_uri: str
- """The target URI to use."""
- _client: Any = PrivateAttr()
- """The parameters to use for queries."""
- query_params: Dict[str, str] = {}
- """The parameters to use for documents."""
- documents_params: Dict[str, str] = {}
-
- def __init__(self, **kwargs: Any):
- super().__init__(**kwargs)
- self._validate_uri()
- try:
- from mlflow.deployments import get_deploy_client
-
- self._client = get_deploy_client(self.target_uri)
- except ImportError as e:
- raise ImportError(
- "Failed to create the client. "
- f"Please run `pip install mlflow{self._mlflow_extras}` to install "
- "required dependencies."
- ) from e
-
- @property
- def _mlflow_extras(self) -> str:
- return "[genai]"
-
- def _validate_uri(self) -> None:
- if self.target_uri == "databricks":
- return
- allowed = ["http", "https", "databricks"]
- if urlparse(self.target_uri).scheme not in allowed:
- raise ValueError(
- f"Invalid target URI: {self.target_uri}. "
- f"The scheme must be one of {allowed}."
- )
-
- def embed(self, texts: List[str], params: Dict[str, str]) -> List[List[float]]:
- embeddings: List[List[float]] = []
- for txt in _chunk(texts, 20):
- resp = self._client.predict(
- endpoint=self.endpoint,
- inputs={"input": txt, **params},
- )
- embeddings.extend(r["embedding"] for r in resp["data"])
- return embeddings
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- return self.embed(texts, params=self.documents_params)
-
- def embed_query(self, text: str) -> List[float]:
- return self.embed([text], params=self.query_params)[0]
-
-
-class MlflowCohereEmbeddings(MlflowEmbeddings):
- """Cohere embedding LLMs in MLflow."""
-
- query_params: Dict[str, str] = {"input_type": "search_query"}
- documents_params: Dict[str, str] = {"input_type": "search_document"}
diff --git a/libs/community/langchain_community/embeddings/mlflow_gateway.py b/libs/community/langchain_community/embeddings/mlflow_gateway.py
deleted file mode 100644
index 9a7a9643fe..0000000000
--- a/libs/community/langchain_community/embeddings/mlflow_gateway.py
+++ /dev/null
@@ -1,79 +0,0 @@
-from __future__ import annotations
-
-import warnings
-from typing import Any, Iterator, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel
-
-
-def _chunk(texts: List[str], size: int) -> Iterator[List[str]]:
- for i in range(0, len(texts), size):
- yield texts[i : i + size]
-
-
-class MlflowAIGatewayEmbeddings(Embeddings, BaseModel):
- """MLflow AI Gateway embeddings.
-
- To use, you should have the ``mlflow[gateway]`` python package installed.
- For more information, see https://mlflow.org/docs/latest/gateway/index.html.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import MlflowAIGatewayEmbeddings
-
- embeddings = MlflowAIGatewayEmbeddings(
- gateway_uri="",
- route=""
- )
- """
-
- route: str
- """The route to use for the MLflow AI Gateway API."""
- gateway_uri: Optional[str] = None
- """The URI for the MLflow AI Gateway API."""
-
- def __init__(self, **kwargs: Any):
- warnings.warn(
- "`MlflowAIGatewayEmbeddings` is deprecated. Use `MlflowEmbeddings` or "
- "`DatabricksEmbeddings` instead.",
- DeprecationWarning,
- )
- try:
- import mlflow.gateway
- except ImportError as e:
- raise ImportError(
- "Could not import `mlflow.gateway` module. "
- "Please install it with `pip install mlflow[gateway]`."
- ) from e
-
- super().__init__(**kwargs)
- if self.gateway_uri:
- mlflow.gateway.set_gateway_uri(self.gateway_uri)
-
- def _query(self, texts: List[str]) -> List[List[float]]:
- try:
- import mlflow.gateway
- except ImportError as e:
- raise ImportError(
- "Could not import `mlflow.gateway` module. "
- "Please install it with `pip install mlflow[gateway]`."
- ) from e
-
- embeddings = []
- for txt in _chunk(texts, 20):
- resp = mlflow.gateway.query(self.route, data={"text": txt})
- # response is List[List[float]]
- if isinstance(resp["embeddings"][0], List):
- embeddings.extend(resp["embeddings"])
- # response is List[float]
- else:
- embeddings.append(resp["embeddings"])
- return embeddings
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- return self._query(texts)
-
- def embed_query(self, text: str) -> List[float]:
- return self._query([text])[0]
diff --git a/libs/community/langchain_community/embeddings/model2vec.py b/libs/community/langchain_community/embeddings/model2vec.py
deleted file mode 100644
index 223f611b0e..0000000000
--- a/libs/community/langchain_community/embeddings/model2vec.py
+++ /dev/null
@@ -1,66 +0,0 @@
-"""Wrapper around model2vec embedding models."""
-
-from typing import List
-
-from langchain_core.embeddings import Embeddings
-
-
-class Model2vecEmbeddings(Embeddings):
- """Model2Vec embedding models.
-
- Install model2vec first, run 'pip install -U model2vec'.
- The github repository for model2vec is : https://github.com/MinishLab/model2vec
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import Model2vecEmbeddings
-
- embedding = Model2vecEmbeddings("minishlab/potion-base-8M")
- embedding.embed_documents([
- "It's dangerous to go alone!",
- "It's a secret to everybody.",
- ])
- embedding.embed_query(
- "Take this with you."
- )
- """
-
- def __init__(self, model: str):
- """Initialize embeddings.
-
- Args:
- model: Model name.
- """
- try:
- from model2vec import StaticModel
- except ImportError as e:
- raise ImportError(
- "Unable to import model2vec, please install with "
- "`pip install -U model2vec`."
- ) from e
- self._model = StaticModel.from_pretrained(model)
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed documents using the model2vec embeddings model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
-
- return self._model.encode(texts).tolist()
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using the model2vec embeddings model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
-
- return self._model.encode(text).tolist()
diff --git a/libs/community/langchain_community/embeddings/modelscope_hub.py b/libs/community/langchain_community/embeddings/modelscope_hub.py
deleted file mode 100644
index e200244c55..0000000000
--- a/libs/community/langchain_community/embeddings/modelscope_hub.py
+++ /dev/null
@@ -1,70 +0,0 @@
-from typing import Any, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict
-
-
-class ModelScopeEmbeddings(BaseModel, Embeddings):
- """ModelScopeHub embedding models.
-
- To use, you should have the ``modelscope`` python package installed.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import ModelScopeEmbeddings
- model_id = "damo/nlp_corom_sentence-embedding_english-base"
- embed = ModelScopeEmbeddings(model_id=model_id, model_revision="v1.0.0")
- """
-
- embed: Any = None
- model_id: str = "damo/nlp_corom_sentence-embedding_english-base"
- """Model name to use."""
- model_revision: Optional[str] = None
-
- def __init__(self, **kwargs: Any):
- """Initialize the modelscope"""
- super().__init__(**kwargs)
- try:
- from modelscope.pipelines import pipeline
- from modelscope.utils.constant import Tasks
- except ImportError as e:
- raise ImportError(
- "Could not import some python packages."
- "Please install it with `pip install modelscope`."
- ) from e
- self.embed = pipeline(
- Tasks.sentence_embedding,
- model=self.model_id,
- model_revision=self.model_revision,
- )
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Compute doc embeddings using a modelscope embedding model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- texts = list(map(lambda x: x.replace("\n", " "), texts))
- inputs = {"source_sentence": texts}
- embeddings = self.embed(input=inputs)["text_embedding"]
- return embeddings.tolist()
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using a modelscope embedding model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- text = text.replace("\n", " ")
- inputs = {"source_sentence": [text]}
- embedding = self.embed(input=inputs)["text_embedding"][0]
- return embedding.tolist()
diff --git a/libs/community/langchain_community/embeddings/mosaicml.py b/libs/community/langchain_community/embeddings/mosaicml.py
deleted file mode 100644
index cf5f8646b3..0000000000
--- a/libs/community/langchain_community/embeddings/mosaicml.py
+++ /dev/null
@@ -1,147 +0,0 @@
-from typing import Any, Dict, List, Mapping, Optional, Tuple
-
-import requests
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import BaseModel, ConfigDict, model_validator
-
-
-class MosaicMLInstructorEmbeddings(BaseModel, Embeddings):
- """MosaicML embedding service.
-
- To use, you should have the
- environment variable ``MOSAICML_API_TOKEN`` set with your API token, or pass
- it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import MosaicMLInstructorEmbeddings
- endpoint_url = (
- "https://models.hosted-on.mosaicml.hosting/instructor-large/v1/predict"
- )
- mosaic_llm = MosaicMLInstructorEmbeddings(
- endpoint_url=endpoint_url,
- mosaicml_api_token="my-api-key"
- )
- """
-
- endpoint_url: str = (
- "https://models.hosted-on.mosaicml.hosting/instructor-xl/v1/predict"
- )
- """Endpoint URL to use."""
- embed_instruction: str = "Represent the document for retrieval: "
- """Instruction used to embed documents."""
- query_instruction: str = (
- "Represent the question for retrieving supporting documents: "
- )
- """Instruction used to embed the query."""
- retry_sleep: float = 1.0
- """How long to try sleeping for if a rate limit is encountered"""
-
- mosaicml_api_token: Optional[str] = None
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and python package exists in environment."""
- mosaicml_api_token = get_from_dict_or_env(
- values, "mosaicml_api_token", "MOSAICML_API_TOKEN"
- )
- values["mosaicml_api_token"] = mosaicml_api_token
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {"endpoint_url": self.endpoint_url}
-
- def _embed(
- self, input: List[Tuple[str, str]], is_retry: bool = False
- ) -> List[List[float]]:
- payload = {"inputs": input}
-
- # HTTP headers for authorization
- headers = {
- "Authorization": f"{self.mosaicml_api_token}",
- "Content-Type": "application/json",
- }
-
- # send request
- try:
- response = requests.post(self.endpoint_url, headers=headers, json=payload)
- except requests.exceptions.RequestException as e:
- raise ValueError(f"Error raised by inference endpoint: {e}")
-
- try:
- if response.status_code == 429:
- if not is_retry:
- import time
-
- time.sleep(self.retry_sleep)
-
- return self._embed(input, is_retry=True)
-
- raise ValueError(
- f"Error raised by inference API: rate limit exceeded.\nResponse: "
- f"{response.text}"
- )
-
- parsed_response = response.json()
-
- # The inference API has changed a couple of times, so we add some handling
- # to be robust to multiple response formats.
- if isinstance(parsed_response, dict):
- output_keys = ["data", "output", "outputs"]
- for key in output_keys:
- if key in parsed_response:
- output_item = parsed_response[key]
- break
- else:
- raise ValueError(
- f"No key data or output in response: {parsed_response}"
- )
-
- if isinstance(output_item, list) and isinstance(output_item[0], list):
- embeddings = output_item
- else:
- embeddings = [output_item]
- else:
- raise ValueError(f"Unexpected response type: {parsed_response}")
-
- except requests.exceptions.JSONDecodeError as e:
- raise ValueError(
- f"Error raised by inference API: {e}.\nResponse: {response.text}"
- )
-
- return embeddings
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed documents using a MosaicML deployed instructor embedding model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- instruction_pairs = [(self.embed_instruction, text) for text in texts]
- embeddings = self._embed(instruction_pairs)
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using a MosaicML deployed instructor embedding model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- instruction_pair = (self.query_instruction, text)
- embedding = self._embed([instruction_pair])[0]
- return embedding
diff --git a/libs/community/langchain_community/embeddings/naver.py b/libs/community/langchain_community/embeddings/naver.py
deleted file mode 100644
index ce20130e52..0000000000
--- a/libs/community/langchain_community/embeddings/naver.py
+++ /dev/null
@@ -1,236 +0,0 @@
-import logging
-from typing import Any, Dict, List, Optional, cast
-
-import httpx
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import convert_to_secret_str, get_from_env
-from pydantic import (
- AliasChoices,
- BaseModel,
- ConfigDict,
- Field,
- SecretStr,
- model_validator,
-)
-from typing_extensions import Self
-
-_DEFAULT_BASE_URL = "https://clovastudio.apigw.ntruss.com"
-_DEFAULT_BASE_URL_ON_NEW_API_KEY = "https://clovastudio.stream.ntruss.com"
-
-logger = logging.getLogger(__name__)
-
-
-def _raise_on_error(response: httpx.Response) -> None:
- """Raise an error if the response is an error."""
- if httpx.codes.is_error(response.status_code):
- error_message = response.read().decode("utf-8")
- raise httpx.HTTPStatusError(
- f"Error response {response.status_code} "
- f"while fetching {response.url}: {error_message}",
- request=response.request,
- response=response,
- )
-
-
-async def _araise_on_error(response: httpx.Response) -> None:
- """Raise an error if the response is an error."""
- if httpx.codes.is_error(response.status_code):
- error_message = (await response.aread()).decode("utf-8")
- raise httpx.HTTPStatusError(
- f"Error response {response.status_code} "
- f"while fetching {response.url}: {error_message}",
- request=response.request,
- response=response,
- )
-
-
-class ClovaXEmbeddings(BaseModel, Embeddings):
- """`NCP ClovaStudio` Embedding API.
-
- following environment variables set or passed in constructor in lower case:
- - ``NCP_CLOVASTUDIO_API_KEY``
- - ``NCP_APIGW_API_KEY``
- - ``NCP_CLOVASTUDIO_APP_ID``
-
- Example:
- .. code-block:: python
-
- from langchain_community import ClovaXEmbeddings
-
- model = ClovaXEmbeddings(model="clir-emb-dolphin")
- output = embedding.embed_documents(documents)
- """ # noqa: E501
-
- client: Optional[httpx.Client] = Field(default=None) #: :meta private:
- async_client: Optional[httpx.AsyncClient] = Field(default=None) #: :meta private:
-
- ncp_clovastudio_api_key: Optional[SecretStr] = Field(default=None, alias="api_key")
- """Automatically inferred from env are `NCP_CLOVASTUDIO_API_KEY` if not provided."""
-
- ncp_apigw_api_key: Optional[SecretStr] = Field(default=None, alias="apigw_api_key")
- """Automatically inferred from env are `NCP_APIGW_API_KEY` if not provided."""
-
- base_url: Optional[str] = Field(default=None, alias="base_url")
- """
- Automatically inferred from env are `NCP_CLOVASTUDIO_API_BASE_URL` if not provided.
- """
-
- app_id: Optional[str] = Field(default=None)
- service_app: bool = Field(
- default=False,
- description="false: use testapp, true: use service app on NCP Clova Studio",
- )
- model_name: str = Field(
- default="clir-emb-dolphin",
- validation_alias=AliasChoices("model_name", "model"),
- description="NCP ClovaStudio embedding model name",
- )
-
- timeout: int = Field(gt=0, default=60)
-
- model_config = ConfigDict(arbitrary_types_allowed=True, protected_namespaces=())
-
- @property
- def lc_secrets(self) -> Dict[str, str]:
- if not self._is_new_api_key():
- return {
- "ncp_clovastudio_api_key": "NCP_CLOVASTUDIO_API_KEY",
- }
- else:
- return {
- "ncp_clovastudio_api_key": "NCP_CLOVASTUDIO_API_KEY",
- "ncp_apigw_api_key": "NCP_APIGW_API_KEY",
- }
-
- @property
- def _api_url(self) -> str:
- """GET embedding api url"""
- app_type = "serviceapp" if self.service_app else "testapp"
- model_name = self.model_name if self.model_name != "bge-m3" else "v2"
- if self._is_new_api_key():
- return f"{self.base_url}/{app_type}/v1/api-tools/embedding/{model_name}"
- else:
- return (
- f"{self.base_url}/{app_type}"
- f"/v1/api-tools/embedding/{model_name}/{self.app_id}"
- )
-
- @model_validator(mode="after")
- def validate_model_after(self) -> Self:
- if not self.ncp_clovastudio_api_key:
- self.ncp_clovastudio_api_key = convert_to_secret_str(
- get_from_env("ncp_clovastudio_api_key", "NCP_CLOVASTUDIO_API_KEY")
- )
-
- if self._is_new_api_key():
- self._init_fields_on_new_api_key()
- else:
- self._init_fields_on_old_api_key()
-
- if not self.base_url:
- raise ValueError("base_url dose not exist.")
-
- if not self.client:
- self.client = httpx.Client(
- base_url=self.base_url,
- headers=self.default_headers(),
- timeout=self.timeout,
- )
-
- if not self.async_client and self.base_url:
- self.async_client = httpx.AsyncClient(
- base_url=self.base_url,
- headers=self.default_headers(),
- timeout=self.timeout,
- )
-
- return self
-
- def _is_new_api_key(self) -> bool:
- if self.ncp_clovastudio_api_key:
- return self.ncp_clovastudio_api_key.get_secret_value().startswith("nv-")
- else:
- return False
-
- def _init_fields_on_new_api_key(self) -> None:
- if not self.base_url:
- self.base_url = get_from_env(
- "base_url",
- "NCP_CLOVASTUDIO_API_BASE_URL",
- _DEFAULT_BASE_URL_ON_NEW_API_KEY,
- )
-
- def _init_fields_on_old_api_key(self) -> None:
- if not self.ncp_apigw_api_key:
- self.ncp_apigw_api_key = convert_to_secret_str(
- get_from_env("ncp_apigw_api_key", "NCP_APIGW_API_KEY", "")
- )
- if not self.base_url:
- self.base_url = get_from_env(
- "base_url", "NCP_CLOVASTUDIO_API_BASE_URL", _DEFAULT_BASE_URL
- )
- if not self.app_id:
- self.app_id = get_from_env("app_id", "NCP_CLOVASTUDIO_APP_ID")
-
- def default_headers(self) -> Dict[str, Any]:
- headers = {
- "Content-Type": "application/json",
- "Accept": "application/json",
- }
-
- clovastudio_api_key = (
- self.ncp_clovastudio_api_key.get_secret_value()
- if self.ncp_clovastudio_api_key
- else None
- )
-
- if self._is_new_api_key():
- ### headers on new api key
- headers["Authorization"] = f"Bearer {clovastudio_api_key}"
- else:
- ### headers on old api key
- if clovastudio_api_key:
- headers["X-NCP-CLOVASTUDIO-API-KEY"] = clovastudio_api_key
-
- apigw_api_key = (
- self.ncp_apigw_api_key.get_secret_value()
- if self.ncp_apigw_api_key
- else None
- )
- if apigw_api_key:
- headers["X-NCP-APIGW-API-KEY"] = apigw_api_key
-
- return headers
-
- def _embed_text(self, text: str) -> List[float]:
- payload = {"text": text}
- client = cast(httpx.Client, self.client)
- response = client.post(url=self._api_url, json=payload)
- _raise_on_error(response)
- return response.json()["result"]["embedding"]
-
- async def _aembed_text(self, text: str) -> List[float]:
- payload = {"text": text}
- async_client = cast(httpx.AsyncClient, self.async_client)
- response = await async_client.post(url=self._api_url, json=payload)
- await _araise_on_error(response)
- return response.json()["result"]["embedding"]
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- embeddings = []
- for text in texts:
- embeddings.append(self._embed_text(text))
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- return self._embed_text(text)
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- embeddings = []
- for text in texts:
- embedding = await self._aembed_text(text)
- embeddings.append(embedding)
- return embeddings
-
- async def aembed_query(self, text: str) -> List[float]:
- return await self._aembed_text(text)
diff --git a/libs/community/langchain_community/embeddings/nemo.py b/libs/community/langchain_community/embeddings/nemo.py
deleted file mode 100644
index fb71bd5e3c..0000000000
--- a/libs/community/langchain_community/embeddings/nemo.py
+++ /dev/null
@@ -1,190 +0,0 @@
-from __future__ import annotations
-
-import asyncio
-import json
-from typing import Any, Dict, List, Optional
-
-import aiohttp
-import requests
-from langchain_core._api.deprecation import deprecated
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import pre_init
-from pydantic import BaseModel
-
-
-def is_endpoint_live(url: str, headers: Optional[dict], payload: Any) -> bool:
- """
- Check if an endpoint is live by sending a GET request to the specified URL.
-
- Args:
- url (str): The URL of the endpoint to check.
-
- Returns:
- bool: True if the endpoint is live (status code 200), False otherwise.
-
- Raises:
- Exception: If the endpoint returns a non-successful status code or if there is
- an error querying the endpoint.
- """
- try:
- response = requests.request("POST", url, headers=headers, data=payload)
-
- # Check if the status code is 200 (OK)
- if response.status_code == 200:
- return True
- else:
- # Raise an exception if the status code is not 200
- raise Exception(
- f"Endpoint returned a non-successful status code: "
- f"{response.status_code}"
- )
- except requests.exceptions.RequestException as e:
- # Handle any exceptions (e.g., connection errors)
- raise Exception(f"Error querying the endpoint: {e}")
-
-
-@deprecated(
- since="0.0.37",
- removal="1.0.0",
- message=(
- "Directly instantiating a NeMoEmbeddings from langchain-community is "
- "deprecated. Please use langchain-nvidia-ai-endpoints NVIDIAEmbeddings "
- "interface."
- ),
-)
-class NeMoEmbeddings(BaseModel, Embeddings):
- """NeMo embedding models."""
-
- batch_size: int = 16
- model: str = "NV-Embed-QA-003"
- api_endpoint_url: str = "http://localhost:8088/v1/embeddings"
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that the end point is alive using the values that are provided."""
-
- url = values["api_endpoint_url"]
- model = values["model"]
-
- # Optional: A minimal test payload and headers required by the endpoint
- headers = {"Content-Type": "application/json"}
- payload = json.dumps(
- {
- "input": "Hello World",
- "model": model,
- "input_type": "query",
- }
- )
-
- is_endpoint_live(url, headers, payload)
-
- return values
-
- async def _aembedding_func(
- self, session: Any, text: str, input_type: str
- ) -> List[float]:
- """Async call out to embedding endpoint.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
-
- headers = {"Content-Type": "application/json"}
-
- async with session.post(
- self.api_endpoint_url,
- json={"input": text, "model": self.model, "input_type": input_type},
- headers=headers,
- ) as response:
- response.raise_for_status()
- answer = await response.text()
- answer = json.loads(answer)
- return answer["data"][0]["embedding"]
-
- def _embedding_func(self, text: str, input_type: str) -> List[float]:
- """Call out to Cohere's embedding endpoint.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
-
- payload = json.dumps(
- {
- "input": text,
- "model": self.model,
- "input_type": input_type,
- }
- )
- headers = {"Content-Type": "application/json"}
-
- response = requests.request(
- "POST", self.api_endpoint_url, headers=headers, data=payload
- )
- response_json = json.loads(response.text)
- embedding = response_json["data"][0]["embedding"]
-
- return embedding
-
- def embed_documents(self, documents: List[str]) -> List[List[float]]:
- """Embed a list of document texts.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- return [self._embedding_func(text, input_type="passage") for text in documents]
-
- def embed_query(self, text: str) -> List[float]:
- return self._embedding_func(text, input_type="query")
-
- async def aembed_query(self, text: str) -> List[float]:
- """Call out to NeMo's embedding endpoint async for embedding query text.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embedding for the text.
- """
-
- async with aiohttp.ClientSession() as session:
- embedding = await self._aembedding_func(session, text, "passage")
- return embedding
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- """Call out to NeMo's embedding endpoint async for embedding search docs.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- embeddings = []
-
- async with aiohttp.ClientSession() as session:
- for batch in range(0, len(texts), self.batch_size):
- text_batch = texts[batch : batch + self.batch_size]
-
- for text in text_batch:
- # Create tasks for all texts in the batch
- tasks = [
- self._aembedding_func(session, text, "passage")
- for text in text_batch
- ]
-
- # Run all tasks concurrently
- batch_results = await asyncio.gather(*tasks)
-
- # Extend the embeddings list with results from this batch
- embeddings.extend(batch_results)
-
- return embeddings
diff --git a/libs/community/langchain_community/embeddings/nlpcloud.py b/libs/community/langchain_community/embeddings/nlpcloud.py
deleted file mode 100644
index 7e13f9cbba..0000000000
--- a/libs/community/langchain_community/embeddings/nlpcloud.py
+++ /dev/null
@@ -1,75 +0,0 @@
-from typing import Any, Dict, List
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict
-
-
-class NLPCloudEmbeddings(BaseModel, Embeddings):
- """NLP Cloud embedding models.
-
- To use, you should have the nlpcloud python package installed
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import NLPCloudEmbeddings
-
- embeddings = NLPCloudEmbeddings()
- """
-
- model_name: str # Define model_name as a class attribute
- gpu: bool # Define gpu as a class attribute
- client: Any #: :meta private:
-
- model_config = ConfigDict(protected_namespaces=())
-
- def __init__(
- self,
- model_name: str = "paraphrase-multilingual-mpnet-base-v2",
- gpu: bool = False,
- **kwargs: Any,
- ) -> None:
- super().__init__(model_name=model_name, gpu=gpu, **kwargs)
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- nlpcloud_api_key = get_from_dict_or_env(
- values, "nlpcloud_api_key", "NLPCLOUD_API_KEY"
- )
- try:
- import nlpcloud
-
- values["client"] = nlpcloud.Client(
- values["model_name"], nlpcloud_api_key, gpu=values["gpu"], lang="en"
- )
- except ImportError:
- raise ImportError(
- "Could not import nlpcloud python package. "
- "Please install it with `pip install nlpcloud`."
- )
- return values
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed a list of documents using NLP Cloud.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
-
- return self.client.embeddings(texts)["embeddings"]
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using NLP Cloud.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self.client.embeddings([text])["embeddings"][0]
diff --git a/libs/community/langchain_community/embeddings/oci_generative_ai.py b/libs/community/langchain_community/embeddings/oci_generative_ai.py
deleted file mode 100644
index 4940f1747c..0000000000
--- a/libs/community/langchain_community/embeddings/oci_generative_ai.py
+++ /dev/null
@@ -1,227 +0,0 @@
-from enum import Enum
-from typing import Any, Dict, Iterator, List, Mapping, Optional
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import pre_init
-from pydantic import BaseModel, ConfigDict
-
-CUSTOM_ENDPOINT_PREFIX = "ocid1.generativeaiendpoint"
-
-
-class OCIAuthType(Enum):
- """OCI authentication types as enumerator."""
-
- API_KEY = 1
- SECURITY_TOKEN = 2
- INSTANCE_PRINCIPAL = 3
- RESOURCE_PRINCIPAL = 4
-
-
-class OCIGenAIEmbeddings(BaseModel, Embeddings):
- """OCI embedding models.
-
- To authenticate, the OCI client uses the methods described in
- https://docs.oracle.com/en-us/iaas/Content/API/Concepts/sdk_authentication_methods.htm
-
- The authentifcation method is passed through auth_type and should be one of:
- API_KEY (default), SECURITY_TOKEN, INSTANCE_PRINCIPLE, RESOURCE_PRINCIPLE
-
- Make sure you have the required policies (profile/roles) to
- access the OCI Generative AI service. If a specific config profile is used,
- you must pass the name of the profile (~/.oci/config) through auth_profile.
- If a specific config file location is used, you must pass
- the file location where profile name configs present
- through auth_file_location
-
- To use, you must provide the compartment id
- along with the endpoint url, and model id
- as named parameters to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain.embeddings import OCIGenAIEmbeddings
-
- embeddings = OCIGenAIEmbeddings(
- model_id="MY_EMBEDDING_MODEL",
- service_endpoint="https://inference.generativeai.us-chicago-1.oci.oraclecloud.com",
- compartment_id="MY_OCID"
- )
- """
-
- client: Any = None #: :meta private:
-
- service_models: Any = None #: :meta private:
-
- auth_type: Optional[str] = "API_KEY"
- """Authentication type, could be
-
- API_KEY,
- SECURITY_TOKEN,
- INSTANCE_PRINCIPLE,
- RESOURCE_PRINCIPLE
-
- If not specified, API_KEY will be used
- """
-
- auth_profile: Optional[str] = "DEFAULT"
- """The name of the profile in ~/.oci/config
- If not specified , DEFAULT will be used
- """
-
- auth_file_location: Optional[str] = "~/.oci/config"
- """Path to the config file.
- If not specified, ~/.oci/config will be used
- """
-
- model_id: Optional[str] = None
- """Id of the model to call, e.g., cohere.embed-english-light-v2.0"""
-
- model_kwargs: Optional[Dict] = None
- """Keyword arguments to pass to the model"""
-
- service_endpoint: Optional[str] = None
- """service endpoint url"""
-
- compartment_id: Optional[str] = None
- """OCID of compartment"""
-
- truncate: Optional[str] = "END"
- """Truncate embeddings that are too long from start or end ("NONE"|"START"|"END")"""
-
- batch_size: int = 96
- """Batch size of OCI GenAI embedding requests. OCI GenAI may handle up to 96 texts
- per request"""
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict: # pylint: disable=no-self-argument
- """Validate that OCI config and python package exists in environment."""
-
- # Skip creating new client if passed in constructor
- if values["client"] is not None:
- return values
-
- try:
- import oci
-
- client_kwargs = {
- "config": {},
- "signer": None,
- "service_endpoint": values["service_endpoint"],
- "retry_strategy": oci.retry.DEFAULT_RETRY_STRATEGY,
- "timeout": (10, 240), # default timeout config for OCI Gen AI service
- }
-
- if values["auth_type"] == OCIAuthType(1).name:
- client_kwargs["config"] = oci.config.from_file(
- file_location=values["auth_file_location"],
- profile_name=values["auth_profile"],
- )
- client_kwargs.pop("signer", None)
- elif values["auth_type"] == OCIAuthType(2).name:
-
- def make_security_token_signer(oci_config): # type: ignore[no-untyped-def]
- pk = oci.signer.load_private_key_from_file(
- oci_config.get("key_file"), None
- )
- with open(
- oci_config.get("security_token_file"), encoding="utf-8"
- ) as f:
- st_string = f.read()
- return oci.auth.signers.SecurityTokenSigner(st_string, pk)
-
- client_kwargs["config"] = oci.config.from_file(
- file_location=values["auth_file_location"],
- profile_name=values["auth_profile"],
- )
- client_kwargs["signer"] = make_security_token_signer(
- oci_config=client_kwargs["config"]
- )
- elif values["auth_type"] == OCIAuthType(3).name:
- client_kwargs["signer"] = (
- oci.auth.signers.InstancePrincipalsSecurityTokenSigner()
- )
- elif values["auth_type"] == OCIAuthType(4).name:
- client_kwargs["signer"] = (
- oci.auth.signers.get_resource_principals_signer()
- )
- else:
- raise ValueError("Please provide valid value to auth_type")
-
- values["client"] = oci.generative_ai_inference.GenerativeAiInferenceClient(
- **client_kwargs
- )
-
- except ImportError as ex:
- raise ImportError(
- "Could not import oci python package. "
- "Please make sure you have the oci package installed."
- ) from ex
- except Exception as e:
- raise ValueError(
- """Could not authenticate with OCI client.
- If INSTANCE_PRINCIPAL or RESOURCE_PRINCIPAL is used,
- please check the specified
- auth_profile, auth_file_location and auth_type are valid.""",
- e,
- ) from e
-
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- **{"model_kwargs": _model_kwargs},
- }
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Call out to OCIGenAI's embedding endpoint.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- from oci.generative_ai_inference import models
-
- if not self.model_id:
- raise ValueError("Model ID is required to embed documents")
-
- if self.model_id.startswith(CUSTOM_ENDPOINT_PREFIX):
- serving_mode = models.DedicatedServingMode(endpoint_id=self.model_id)
- else:
- serving_mode = models.OnDemandServingMode(model_id=self.model_id)
-
- embeddings = []
-
- def split_texts() -> Iterator[List[str]]:
- for i in range(0, len(texts), self.batch_size):
- yield texts[i : i + self.batch_size]
-
- for chunk in split_texts():
- invocation_obj = models.EmbedTextDetails(
- serving_mode=serving_mode,
- compartment_id=self.compartment_id,
- truncate=self.truncate,
- inputs=chunk,
- )
- response = self.client.embed_text(invocation_obj)
- embeddings.extend(response.data.embeddings)
-
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to OCIGenAI's embedding endpoint.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self.embed_documents([text])[0]
diff --git a/libs/community/langchain_community/embeddings/octoai_embeddings.py b/libs/community/langchain_community/embeddings/octoai_embeddings.py
deleted file mode 100644
index cd10033e38..0000000000
--- a/libs/community/langchain_community/embeddings/octoai_embeddings.py
+++ /dev/null
@@ -1,86 +0,0 @@
-from typing import Dict, Optional
-
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import Field, SecretStr
-
-from langchain_community.embeddings.openai import OpenAIEmbeddings
-from langchain_community.utils.openai import is_openai_v1
-
-DEFAULT_API_BASE = "https://text.octoai.run/v1/"
-DEFAULT_MODEL = "thenlper/gte-large"
-
-
-class OctoAIEmbeddings(OpenAIEmbeddings):
- """OctoAI Compute Service embedding models.
-
- See https://octo.ai/ for information about OctoAI.
-
- To use, you should have the ``openai`` python package installed and the
- environment variable ``OCTOAI_API_TOKEN`` set with your API token.
- Alternatively, you can use the octoai_api_token keyword argument.
- """
-
- octoai_api_token: Optional[SecretStr] = Field(default=None)
- """OctoAI Endpoints API keys."""
- endpoint_url: str = Field(default=DEFAULT_API_BASE)
- """Base URL path for API requests."""
- model: str = Field(default=DEFAULT_MODEL)
- """Model name to use."""
- tiktoken_enabled: bool = False
- """Set this to False for non-OpenAI implementations of the embeddings API"""
-
- @property
- def _llm_type(self) -> str:
- """Return type of embeddings model."""
- return "octoai-embeddings"
-
- @property
- def lc_secrets(self) -> Dict[str, str]:
- return {"octoai_api_token": "OCTOAI_API_TOKEN"}
-
- @pre_init
- def validate_environment(cls, values: dict) -> dict:
- """Validate that api key and python package exists in environment."""
- values["endpoint_url"] = get_from_dict_or_env(
- values,
- "endpoint_url",
- "ENDPOINT_URL",
- default=DEFAULT_API_BASE,
- )
- values["octoai_api_token"] = convert_to_secret_str(
- get_from_dict_or_env(values, "octoai_api_token", "OCTOAI_API_TOKEN")
- )
- values["model"] = get_from_dict_or_env(
- values,
- "model",
- "MODEL",
- default=DEFAULT_MODEL,
- )
-
- try:
- import openai
-
- if is_openai_v1():
- client_params = {
- "api_key": values["octoai_api_token"].get_secret_value(),
- "base_url": values["endpoint_url"],
- }
- if not values.get("client"):
- values["client"] = openai.OpenAI(**client_params).embeddings
- if not values.get("async_client"):
- values["async_client"] = openai.AsyncOpenAI(
- **client_params
- ).embeddings
- else:
- values["openai_api_base"] = values["endpoint_url"]
- values["openai_api_key"] = values["octoai_api_token"].get_secret_value()
- values["client"] = openai.Embedding
- values["async_client"] = openai.Embedding
-
- except ImportError:
- raise ImportError(
- "Could not import openai python package. "
- "Please install it with `pip install openai`."
- )
-
- return values
diff --git a/libs/community/langchain_community/embeddings/ollama.py b/libs/community/langchain_community/embeddings/ollama.py
deleted file mode 100644
index ddec6fe39b..0000000000
--- a/libs/community/langchain_community/embeddings/ollama.py
+++ /dev/null
@@ -1,228 +0,0 @@
-import logging
-from typing import Any, Dict, List, Mapping, Optional
-
-import requests
-from langchain_core._api.deprecation import deprecated
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict
-
-logger = logging.getLogger(__name__)
-
-
-@deprecated(
- since="0.3.1",
- removal="1.0.0",
- alternative_import="langchain_ollama.OllamaEmbeddings",
-)
-class OllamaEmbeddings(BaseModel, Embeddings):
- """Ollama locally runs large language models.
-
- To use, follow the instructions at https://ollama.ai/.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import OllamaEmbeddings
- ollama_emb = OllamaEmbeddings(
- model="llama:7b",
- )
- r1 = ollama_emb.embed_documents(
- [
- "Alpha is the first letter of Greek alphabet",
- "Beta is the second letter of Greek alphabet",
- ]
- )
- r2 = ollama_emb.embed_query(
- "What is the second letter of Greek alphabet"
- )
-
- """
-
- base_url: str = "http://localhost:11434"
- """Base url the model is hosted under."""
- model: str = "llama2"
- """Model name to use."""
-
- embed_instruction: str = "passage: "
- """Instruction used to embed documents."""
- query_instruction: str = "query: "
- """Instruction used to embed the query."""
-
- mirostat: Optional[int] = None
- """Enable Mirostat sampling for controlling perplexity.
- (default: 0, 0 = disabled, 1 = Mirostat, 2 = Mirostat 2.0)"""
-
- mirostat_eta: Optional[float] = None
- """Influences how quickly the algorithm responds to feedback
- from the generated text. A lower learning rate will result in
- slower adjustments, while a higher learning rate will make
- the algorithm more responsive. (Default: 0.1)"""
-
- mirostat_tau: Optional[float] = None
- """Controls the balance between coherence and diversity
- of the output. A lower value will result in more focused and
- coherent text. (Default: 5.0)"""
-
- num_ctx: Optional[int] = None
- """Sets the size of the context window used to generate the
- next token. (Default: 2048) """
-
- num_gpu: Optional[int] = None
- """The number of GPUs to use. On macOS it defaults to 1 to
- enable metal support, 0 to disable."""
-
- num_thread: Optional[int] = None
- """Sets the number of threads to use during computation.
- By default, Ollama will detect this for optimal performance.
- It is recommended to set this value to the number of physical
- CPU cores your system has (as opposed to the logical number of cores)."""
-
- repeat_last_n: Optional[int] = None
- """Sets how far back for the model to look back to prevent
- repetition. (Default: 64, 0 = disabled, -1 = num_ctx)"""
-
- repeat_penalty: Optional[float] = None
- """Sets how strongly to penalize repetitions. A higher value (e.g., 1.5)
- will penalize repetitions more strongly, while a lower value (e.g., 0.9)
- will be more lenient. (Default: 1.1)"""
-
- temperature: Optional[float] = None
- """The temperature of the model. Increasing the temperature will
- make the model answer more creatively. (Default: 0.8)"""
-
- stop: Optional[List[str]] = None
- """Sets the stop tokens to use."""
-
- tfs_z: Optional[float] = None
- """Tail free sampling is used to reduce the impact of less probable
- tokens from the output. A higher value (e.g., 2.0) will reduce the
- impact more, while a value of 1.0 disables this setting. (default: 1)"""
-
- top_k: Optional[int] = None
- """Reduces the probability of generating nonsense. A higher value (e.g. 100)
- will give more diverse answers, while a lower value (e.g. 10)
- will be more conservative. (Default: 40)"""
-
- top_p: Optional[float] = None
- """Works together with top-k. A higher value (e.g., 0.95) will lead
- to more diverse text, while a lower value (e.g., 0.5) will
- generate more focused and conservative text. (Default: 0.9)"""
-
- show_progress: bool = False
- """Whether to show a tqdm progress bar. Must have `tqdm` installed."""
-
- headers: Optional[dict] = None
- """Additional headers to pass to endpoint (e.g. Authorization, Referer).
- This is useful when Ollama is hosted on cloud services that require
- tokens for authentication.
- """
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling Ollama."""
- return {
- "model": self.model,
- "options": {
- "mirostat": self.mirostat,
- "mirostat_eta": self.mirostat_eta,
- "mirostat_tau": self.mirostat_tau,
- "num_ctx": self.num_ctx,
- "num_gpu": self.num_gpu,
- "num_thread": self.num_thread,
- "repeat_last_n": self.repeat_last_n,
- "repeat_penalty": self.repeat_penalty,
- "temperature": self.temperature,
- "stop": self.stop,
- "tfs_z": self.tfs_z,
- "top_k": self.top_k,
- "top_p": self.top_p,
- },
- }
-
- model_kwargs: Optional[dict] = None
- """Other model keyword args"""
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {**{"model": self.model}, **self._default_params}
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- def _process_emb_response(self, input: str) -> List[float]:
- """Process a response from the API.
-
- Args:
- response: The response from the API.
-
- Returns:
- The response as a dictionary.
- """
- headers = {
- "Content-Type": "application/json",
- **(self.headers or {}),
- }
-
- try:
- res = requests.post(
- f"{self.base_url}/api/embeddings",
- headers=headers,
- json={"model": self.model, "prompt": input, **self._default_params},
- )
- except requests.exceptions.RequestException as e:
- raise ValueError(f"Error raised by inference endpoint: {e}")
-
- if res.status_code != 200:
- raise ValueError(
- "Error raised by inference API HTTP code: %s, %s"
- % (res.status_code, res.text)
- )
- try:
- t = res.json()
- return t["embedding"]
- except requests.exceptions.JSONDecodeError as e:
- raise ValueError(
- f"Error raised by inference API: {e}.\nResponse: {res.text}"
- )
-
- def _embed(self, input: List[str]) -> List[List[float]]:
- if self.show_progress:
- try:
- from tqdm import tqdm
-
- iter_ = tqdm(input, desc="OllamaEmbeddings")
- except ImportError:
- logger.warning(
- "Unable to show progress bar because tqdm could not be imported. "
- "Please install with `pip install tqdm`."
- )
- iter_ = input
- else:
- iter_ = input
- return [self._process_emb_response(prompt) for prompt in iter_]
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed documents using an Ollama deployed embedding model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- instruction_pairs = [f"{self.embed_instruction}{text}" for text in texts]
- embeddings = self._embed(instruction_pairs)
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using a Ollama deployed embedding model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- instruction_pair = f"{self.query_instruction}{text}"
- embedding = self._embed([instruction_pair])[0]
- return embedding
diff --git a/libs/community/langchain_community/embeddings/openai.py b/libs/community/langchain_community/embeddings/openai.py
deleted file mode 100644
index a695ab72ff..0000000000
--- a/libs/community/langchain_community/embeddings/openai.py
+++ /dev/null
@@ -1,716 +0,0 @@
-from __future__ import annotations
-
-import logging
-import os
-import warnings
-from typing import (
- Any,
- Callable,
- Dict,
- List,
- Literal,
- Mapping,
- Optional,
- Sequence,
- Set,
- Tuple,
- Union,
- cast,
-)
-
-import numpy as np
-from langchain_core._api.deprecation import deprecated
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import (
- get_from_dict_or_env,
- get_pydantic_field_names,
- pre_init,
-)
-from pydantic import BaseModel, ConfigDict, Field, model_validator
-from tenacity import (
- AsyncRetrying,
- before_sleep_log,
- retry,
- retry_if_exception_type,
- stop_after_attempt,
- wait_exponential,
-)
-
-from langchain_community.utils.openai import is_openai_v1
-
-logger = logging.getLogger(__name__)
-
-
-def _create_retry_decorator(embeddings: OpenAIEmbeddings) -> Callable[[Any], Any]:
- import openai
-
- # Wait 2^x * 1 second between each retry starting with
- # retry_min_seconds seconds, then up to retry_max_seconds seconds,
- # then retry_max_seconds seconds afterwards
- # retry_min_seconds and retry_max_seconds are optional arguments of
- # OpenAIEmbeddings
- return retry(
- reraise=True,
- stop=stop_after_attempt(embeddings.max_retries),
- wait=wait_exponential(
- multiplier=1,
- min=embeddings.retry_min_seconds,
- max=embeddings.retry_max_seconds,
- ),
- retry=(
- retry_if_exception_type(openai.error.Timeout)
- | retry_if_exception_type(openai.error.APIError)
- | retry_if_exception_type(openai.error.APIConnectionError)
- | retry_if_exception_type(openai.error.RateLimitError)
- | retry_if_exception_type(openai.error.ServiceUnavailableError)
- ),
- before_sleep=before_sleep_log(logger, logging.WARNING),
- )
-
-
-def _async_retry_decorator(embeddings: OpenAIEmbeddings) -> Any:
- import openai
-
- # Wait 2^x * 1 second between each retry starting with
- # retry_min_seconds seconds, then up to retry_max_seconds seconds,
- # then retry_max_seconds seconds afterwards
- # retry_min_seconds and retry_max_seconds are optional arguments of
- # OpenAIEmbeddings
- async_retrying = AsyncRetrying(
- reraise=True,
- stop=stop_after_attempt(embeddings.max_retries),
- wait=wait_exponential(
- multiplier=1,
- min=embeddings.retry_min_seconds,
- max=embeddings.retry_max_seconds,
- ),
- retry=(
- retry_if_exception_type(openai.error.Timeout)
- | retry_if_exception_type(openai.error.APIError)
- | retry_if_exception_type(openai.error.APIConnectionError)
- | retry_if_exception_type(openai.error.RateLimitError)
- | retry_if_exception_type(openai.error.ServiceUnavailableError)
- ),
- before_sleep=before_sleep_log(logger, logging.WARNING),
- )
-
- def wrap(func: Callable) -> Callable:
- async def wrapped_f(*args: Any, **kwargs: Any) -> Callable:
- async for _ in async_retrying:
- return await func(*args, **kwargs)
- raise AssertionError("this is unreachable")
-
- return wrapped_f
-
- return wrap
-
-
-# https://stackoverflow.com/questions/76469415/getting-embeddings-of-length-1-from-langchain-openaiembeddings
-def _check_response(response: dict, skip_empty: bool = False) -> dict:
- if any(len(d["embedding"]) == 1 for d in response["data"]) and not skip_empty:
- import openai
-
- raise openai.error.APIError("OpenAI API returned an empty embedding")
- return response
-
-
-def embed_with_retry(embeddings: OpenAIEmbeddings, **kwargs: Any) -> Any:
- """Use tenacity to retry the embedding call."""
- if is_openai_v1():
- return embeddings.client.create(**kwargs)
- retry_decorator = _create_retry_decorator(embeddings)
-
- @retry_decorator
- def _embed_with_retry(**kwargs: Any) -> Any:
- response = embeddings.client.create(**kwargs)
- return _check_response(response, skip_empty=embeddings.skip_empty)
-
- return _embed_with_retry(**kwargs)
-
-
-async def async_embed_with_retry(embeddings: OpenAIEmbeddings, **kwargs: Any) -> Any:
- """Use tenacity to retry the embedding call."""
-
- if is_openai_v1():
- return await embeddings.async_client.create(**kwargs)
-
- @_async_retry_decorator(embeddings)
- async def _async_embed_with_retry(**kwargs: Any) -> Any:
- response = await embeddings.client.acreate(**kwargs)
- return _check_response(response, skip_empty=embeddings.skip_empty)
-
- return await _async_embed_with_retry(**kwargs)
-
-
-@deprecated(
- since="0.0.9",
- removal="1.0",
- alternative_import="langchain_openai.OpenAIEmbeddings",
-)
-class OpenAIEmbeddings(BaseModel, Embeddings):
- """OpenAI embedding models.
-
- To use, you should have the ``openai`` python package installed, and the
- environment variable ``OPENAI_API_KEY`` set with your API key or pass it
- as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import OpenAIEmbeddings
- openai = OpenAIEmbeddings(openai_api_key="my-api-key")
-
- In order to use the library with Microsoft Azure endpoints, you need to set
- the OPENAI_API_TYPE, OPENAI_API_BASE, OPENAI_API_KEY and OPENAI_API_VERSION.
- The OPENAI_API_TYPE must be set to 'azure' and the others correspond to
- the properties of your endpoint.
- In addition, the deployment name must be passed as the model parameter.
-
- Example:
- .. code-block:: python
-
- import os
-
- os.environ["OPENAI_API_TYPE"] = "azure"
- os.environ["OPENAI_API_BASE"] = "https:// Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = get_pydantic_field_names(cls)
- extra = values.get("model_kwargs", {})
- for field_name in list(values):
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- if field_name not in all_required_field_names:
- warnings.warn(
- f"""WARNING! {field_name} is not default parameter.
- {field_name} was transferred to model_kwargs.
- Please confirm that {field_name} is what you intended."""
- )
- extra[field_name] = values.pop(field_name)
-
- invalid_model_kwargs = all_required_field_names.intersection(extra.keys())
- if invalid_model_kwargs:
- raise ValueError(
- f"Parameters {invalid_model_kwargs} should be specified explicitly. "
- f"Instead they were passed in as part of `model_kwargs` parameter."
- )
-
- values["model_kwargs"] = extra
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- values["openai_api_key"] = get_from_dict_or_env(
- values, "openai_api_key", "OPENAI_API_KEY"
- )
- values["openai_api_base"] = values["openai_api_base"] or os.getenv(
- "OPENAI_API_BASE"
- )
- values["openai_api_type"] = get_from_dict_or_env(
- values,
- "openai_api_type",
- "OPENAI_API_TYPE",
- default="",
- )
- values["openai_proxy"] = get_from_dict_or_env(
- values,
- "openai_proxy",
- "OPENAI_PROXY",
- default="",
- )
- if values["openai_api_type"] in ("azure", "azure_ad", "azuread"):
- default_api_version = "2023-05-15"
- # Azure OpenAI embedding models allow a maximum of 2048
- # texts at a time in each batch
- # See: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#embeddings
- values["chunk_size"] = min(values["chunk_size"], 2048)
- else:
- default_api_version = ""
- values["openai_api_version"] = get_from_dict_or_env(
- values,
- "openai_api_version",
- "OPENAI_API_VERSION",
- default=default_api_version,
- )
- # Check OPENAI_ORGANIZATION for backwards compatibility.
- values["openai_organization"] = (
- values["openai_organization"]
- or os.getenv("OPENAI_ORG_ID")
- or os.getenv("OPENAI_ORGANIZATION")
- )
- try:
- import openai
- except ImportError:
- raise ImportError(
- "Could not import openai python package. "
- "Please install it with `pip install openai`."
- )
- else:
- if is_openai_v1():
- if values["openai_api_type"] in ("azure", "azure_ad", "azuread"):
- warnings.warn(
- "If you have openai>=1.0.0 installed and are using Azure, "
- "please use the `AzureOpenAIEmbeddings` class."
- )
- client_params = {
- "api_key": values["openai_api_key"],
- "organization": values["openai_organization"],
- "base_url": values["openai_api_base"],
- "timeout": values["request_timeout"],
- "max_retries": values["max_retries"],
- "default_headers": values["default_headers"],
- "default_query": values["default_query"],
- "http_client": values["http_client"],
- }
- if not values.get("client"):
- values["client"] = openai.OpenAI(**client_params).embeddings
- if not values.get("async_client"):
- values["async_client"] = openai.AsyncOpenAI(
- **client_params
- ).embeddings
- elif not values.get("client"):
- values["client"] = openai.Embedding
- else:
- pass
- return values
-
- @property
- def _invocation_params(self) -> Dict[str, Any]:
- if is_openai_v1():
- openai_args: Dict = {"model": self.model, **self.model_kwargs}
- else:
- openai_args = {
- "model": self.model,
- "request_timeout": self.request_timeout,
- "headers": self.headers,
- "api_key": self.openai_api_key,
- "organization": self.openai_organization,
- "api_base": self.openai_api_base,
- "api_type": self.openai_api_type,
- "api_version": self.openai_api_version,
- **self.model_kwargs,
- }
- if self.openai_api_type in ("azure", "azure_ad", "azuread"):
- openai_args["engine"] = self.deployment
- # TODO: Look into proxy with openai v1.
- if self.openai_proxy:
- try:
- import openai
- except ImportError:
- raise ImportError(
- "Could not import openai python package. "
- "Please install it with `pip install openai`."
- )
-
- openai.proxy = {
- "http": self.openai_proxy,
- "https": self.openai_proxy,
- }
- return openai_args
-
- # please refer to
- # https://github.com/openai/openai-cookbook/blob/main/examples/Embedding_long_inputs.ipynb
- def _get_len_safe_embeddings(
- self, texts: List[str], *, engine: str, chunk_size: Optional[int] = None
- ) -> List[List[float]]:
- """
- Generate length-safe embeddings for a list of texts.
-
- This method handles tokenization and embedding generation, respecting the
- set embedding context length and chunk size. It supports both tiktoken
- and HuggingFace tokenizer based on the tiktoken_enabled flag.
-
- Args:
- texts (List[str]): A list of texts to embed.
- engine (str): The engine or model to use for embeddings.
- chunk_size (Optional[int]): The size of chunks for processing embeddings.
-
- Returns:
- List[List[float]]: A list of embeddings for each input text.
- """
-
- tokens = []
- indices = []
- model_name = self.tiktoken_model_name or self.model
- _chunk_size = chunk_size or self.chunk_size
-
- # If tiktoken flag set to False
- if not self.tiktoken_enabled:
- try:
- from transformers import AutoTokenizer
- except ImportError:
- raise ImportError(
- "Could not import transformers python package. "
- "This is needed in order to for OpenAIEmbeddings without "
- "`tiktoken`. Please install it with `pip install transformers`. "
- )
-
- tokenizer = AutoTokenizer.from_pretrained(
- pretrained_model_name_or_path=model_name
- )
- for i, text in enumerate(texts):
- # Tokenize the text using HuggingFace transformers
- tokenized = tokenizer.encode(text, add_special_tokens=False)
-
- # Split tokens into chunks respecting the embedding_ctx_length
- for j in range(0, len(tokenized), self.embedding_ctx_length):
- token_chunk = tokenized[j : j + self.embedding_ctx_length]
-
- # Convert token IDs back to a string
- chunk_text = tokenizer.decode(token_chunk)
- tokens.append(chunk_text)
- indices.append(i)
- else:
- try:
- import tiktoken
- except ImportError:
- raise ImportError(
- "Could not import tiktoken python package. "
- "This is needed in order to for OpenAIEmbeddings. "
- "Please install it with `pip install tiktoken`."
- )
-
- try:
- encoding = tiktoken.encoding_for_model(model_name)
- except KeyError:
- logger.warning("Warning: model not found. Using cl100k_base encoding.")
- model = "cl100k_base"
- encoding = tiktoken.get_encoding(model)
- for i, text in enumerate(texts):
- if self.model.endswith("001"):
- # See: https://github.com/openai/openai-python/
- # issues/418#issuecomment-1525939500
- # replace newlines, which can negatively affect performance.
- text = text.replace("\n", " ")
-
- token = encoding.encode(
- text=text,
- allowed_special=self.allowed_special,
- disallowed_special=self.disallowed_special,
- )
-
- # Split tokens into chunks respecting the embedding_ctx_length
- for j in range(0, len(token), self.embedding_ctx_length):
- tokens.append(token[j : j + self.embedding_ctx_length])
- indices.append(i)
-
- if self.show_progress_bar:
- try:
- from tqdm.auto import tqdm
-
- _iter = tqdm(range(0, len(tokens), _chunk_size))
- except ImportError:
- _iter = range(0, len(tokens), _chunk_size)
- else:
- _iter = range(0, len(tokens), _chunk_size)
-
- batched_embeddings: List[List[float]] = []
- for i in _iter:
- response = embed_with_retry(
- self,
- input=tokens[i : i + _chunk_size],
- **self._invocation_params,
- )
- if not isinstance(response, dict):
- response = response.dict()
- batched_embeddings.extend(r["embedding"] for r in response["data"])
-
- results: List[List[List[float]]] = [[] for _ in range(len(texts))]
- num_tokens_in_batch: List[List[int]] = [[] for _ in range(len(texts))]
- for i in range(len(indices)):
- if self.skip_empty and len(batched_embeddings[i]) == 1:
- continue
- results[indices[i]].append(batched_embeddings[i])
- num_tokens_in_batch[indices[i]].append(len(tokens[i]))
-
- embeddings: List[List[float]] = [[] for _ in range(len(texts))]
- for i in range(len(texts)):
- _result = results[i]
- if len(_result) == 0:
- average_embedded = embed_with_retry(
- self,
- input="",
- **self._invocation_params,
- )
- if not isinstance(average_embedded, dict):
- average_embedded = average_embedded.dict()
- average = average_embedded["data"][0]["embedding"]
- else:
- average = np.average(_result, axis=0, weights=num_tokens_in_batch[i])
- embeddings[i] = (average / np.linalg.norm(average)).tolist()
-
- return embeddings
-
- # please refer to
- # https://github.com/openai/openai-cookbook/blob/main/examples/Embedding_long_inputs.ipynb
- async def _aget_len_safe_embeddings(
- self, texts: List[str], *, engine: str, chunk_size: Optional[int] = None
- ) -> List[List[float]]:
- """
- Asynchronously generate length-safe embeddings for a list of texts.
-
- This method handles tokenization and asynchronous embedding generation,
- respecting the set embedding context length and chunk size. It supports both
- `tiktoken` and HuggingFace `tokenizer` based on the tiktoken_enabled flag.
-
- Args:
- texts (List[str]): A list of texts to embed.
- engine (str): The engine or model to use for embeddings.
- chunk_size (Optional[int]): The size of chunks for processing embeddings.
-
- Returns:
- List[List[float]]: A list of embeddings for each input text.
- """
-
- tokens = []
- indices = []
- model_name = self.tiktoken_model_name or self.model
- _chunk_size = chunk_size or self.chunk_size
-
- # If tiktoken flag set to False
- if not self.tiktoken_enabled:
- try:
- from transformers import AutoTokenizer
- except ImportError:
- raise ImportError(
- "Could not import transformers python package. "
- "This is needed in order to for OpenAIEmbeddings without "
- " `tiktoken`. Please install it with `pip install transformers`."
- )
-
- tokenizer = AutoTokenizer.from_pretrained(
- pretrained_model_name_or_path=model_name
- )
- for i, text in enumerate(texts):
- # Tokenize the text using HuggingFace transformers
- tokenized = tokenizer.encode(text, add_special_tokens=False)
-
- # Split tokens into chunks respecting the embedding_ctx_length
- for j in range(0, len(tokenized), self.embedding_ctx_length):
- token_chunk = tokenized[j : j + self.embedding_ctx_length]
-
- # Convert token IDs back to a string
- chunk_text = tokenizer.decode(token_chunk)
- tokens.append(chunk_text)
- indices.append(i)
- else:
- try:
- import tiktoken
- except ImportError:
- raise ImportError(
- "Could not import tiktoken python package. "
- "This is needed in order to for OpenAIEmbeddings. "
- "Please install it with `pip install tiktoken`."
- )
-
- try:
- encoding = tiktoken.encoding_for_model(model_name)
- except KeyError:
- logger.warning("Warning: model not found. Using cl100k_base encoding.")
- model = "cl100k_base"
- encoding = tiktoken.get_encoding(model)
- for i, text in enumerate(texts):
- if self.model.endswith("001"):
- # See: https://github.com/openai/openai-python/
- # issues/418#issuecomment-1525939500
- # replace newlines, which can negatively affect performance.
- text = text.replace("\n", " ")
-
- token = encoding.encode(
- text=text,
- allowed_special=self.allowed_special,
- disallowed_special=self.disallowed_special,
- )
-
- # Split tokens into chunks respecting the embedding_ctx_length
- for j in range(0, len(token), self.embedding_ctx_length):
- tokens.append(token[j : j + self.embedding_ctx_length])
- indices.append(i)
-
- batched_embeddings: List[List[float]] = []
- _chunk_size = chunk_size or self.chunk_size
- for i in range(0, len(tokens), _chunk_size):
- response = await async_embed_with_retry(
- self,
- input=tokens[i : i + _chunk_size],
- **self._invocation_params,
- )
-
- if not isinstance(response, dict):
- response = response.dict()
- batched_embeddings.extend(r["embedding"] for r in response["data"])
-
- results: List[List[List[float]]] = [[] for _ in range(len(texts))]
- num_tokens_in_batch: List[List[int]] = [[] for _ in range(len(texts))]
- for i in range(len(indices)):
- results[indices[i]].append(batched_embeddings[i])
- num_tokens_in_batch[indices[i]].append(len(tokens[i]))
-
- embeddings: List[List[float]] = [[] for _ in range(len(texts))]
- for i in range(len(texts)):
- _result = results[i]
- if len(_result) == 0:
- average_embedded = await async_embed_with_retry(
- self,
- input="",
- **self._invocation_params,
- )
- if not isinstance(average_embedded, dict):
- average_embedded = average_embedded.dict()
- average = average_embedded["data"][0]["embedding"]
- else:
- average = np.average(_result, axis=0, weights=num_tokens_in_batch[i])
- embeddings[i] = (average / np.linalg.norm(average)).tolist()
-
- return embeddings
-
- def embed_documents(
- self, texts: List[str], chunk_size: Optional[int] = 0
- ) -> List[List[float]]:
- """Call out to OpenAI's embedding endpoint for embedding search docs.
-
- Args:
- texts: The list of texts to embed.
- chunk_size: The chunk size of embeddings. If None, will use the chunk size
- specified by the class.
-
- Returns:
- List of embeddings, one for each text.
- """
- # NOTE: to keep things simple, we assume the list may contain texts longer
- # than the maximum context and use length-safe embedding function.
- engine = cast(str, self.deployment)
- return self._get_len_safe_embeddings(
- texts, engine=engine, chunk_size=chunk_size
- )
-
- async def aembed_documents(
- self, texts: List[str], chunk_size: Optional[int] = 0
- ) -> List[List[float]]:
- """Call out to OpenAI's embedding endpoint async for embedding search docs.
-
- Args:
- texts: The list of texts to embed.
- chunk_size: The chunk size of embeddings. If None, will use the chunk size
- specified by the class.
-
- Returns:
- List of embeddings, one for each text.
- """
- # NOTE: to keep things simple, we assume the list may contain texts longer
- # than the maximum context and use length-safe embedding function.
- engine = cast(str, self.deployment)
- return self._get_len_safe_embeddings(
- texts, engine=engine, chunk_size=chunk_size
- )
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to OpenAI's embedding endpoint for embedding query text.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embedding for the text.
- """
- return self.embed_documents([text])[0]
-
- async def aembed_query(self, text: str) -> List[float]:
- """Call out to OpenAI's embedding endpoint async for embedding query text.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embedding for the text.
- """
- embeddings = await self.aembed_documents([text])
- return embeddings[0]
diff --git a/libs/community/langchain_community/embeddings/openvino.py b/libs/community/langchain_community/embeddings/openvino.py
deleted file mode 100644
index 930453247c..0000000000
--- a/libs/community/langchain_community/embeddings/openvino.py
+++ /dev/null
@@ -1,351 +0,0 @@
-from pathlib import Path
-from typing import Any, Dict, List
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict, Field
-
-DEFAULT_QUERY_INSTRUCTION = (
- "Represent the question for retrieving supporting documents: "
-)
-DEFAULT_QUERY_BGE_INSTRUCTION_EN = (
- "Represent this question for searching relevant passages: "
-)
-DEFAULT_QUERY_BGE_INSTRUCTION_ZH = "为这个句子生成表示以用于检索相关文章:"
-
-
-class OpenVINOEmbeddings(BaseModel, Embeddings):
- """OpenVINO embedding models.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import OpenVINOEmbeddings
-
- model_name = "sentence-transformers/all-mpnet-base-v2"
- model_kwargs = {'device': 'CPU'}
- encode_kwargs = {'normalize_embeddings': True}
- ov = OpenVINOEmbeddings(
- model_name_or_path=model_name,
- model_kwargs=model_kwargs,
- encode_kwargs=encode_kwargs
- )
- """
-
- ov_model: Any = None
- """OpenVINO model object."""
- tokenizer: Any = None
- """Tokenizer for embedding model."""
- model_name_or_path: str
- """HuggingFace model id."""
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Keyword arguments to pass to the model."""
- encode_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Keyword arguments to pass when calling the `encode` method of the model."""
- show_progress: bool = False
- """Whether to show a progress bar."""
-
- def __init__(self, **kwargs: Any):
- """Initialize the sentence_transformer."""
- super().__init__(**kwargs)
-
- try:
- from optimum.intel.openvino import OVModelForFeatureExtraction
- except ImportError as e:
- raise ImportError(
- "Could not import optimum-intel python package. "
- "Please install it with: "
- "pip install -U 'optimum[openvino,nncf]'"
- ) from e
-
- try:
- from huggingface_hub import HfApi
- except ImportError as e:
- raise ImportError(
- "Could not import huggingface_hub python package. "
- "Please install it with: "
- "`pip install -U huggingface_hub`."
- ) from e
-
- def require_model_export(
- model_id: str, revision: Any = None, subfolder: Any = None
- ) -> bool:
- model_dir = Path(model_id)
- if subfolder is not None:
- model_dir = model_dir / subfolder
- if model_dir.is_dir():
- return (
- not (model_dir / "openvino_model.xml").exists()
- or not (model_dir / "openvino_model.bin").exists()
- )
- hf_api = HfApi()
- try:
- model_info = hf_api.model_info(model_id, revision=revision or "main")
- normalized_subfolder = (
- None if subfolder is None else Path(subfolder).as_posix()
- )
- model_files = [
- file.rfilename
- for file in model_info.siblings
- if normalized_subfolder is None
- or file.rfilename.startswith(normalized_subfolder)
- ]
- ov_model_path = (
- "openvino_model.xml"
- if subfolder is None
- else f"{normalized_subfolder}/openvino_model.xml"
- )
- return (
- ov_model_path not in model_files
- or ov_model_path.replace(".xml", ".bin") not in model_files
- )
- except Exception:
- return True
-
- if require_model_export(self.model_name_or_path):
- # use remote model
- self.ov_model = OVModelForFeatureExtraction.from_pretrained(
- self.model_name_or_path, export=True, **self.model_kwargs
- )
- else:
- # use local model
- self.ov_model = OVModelForFeatureExtraction.from_pretrained(
- self.model_name_or_path, **self.model_kwargs
- )
-
- try:
- from transformers import AutoTokenizer
- except ImportError as e:
- raise ImportError(
- "Unable to import transformers, please install with "
- "`pip install -U transformers`."
- ) from e
- self.tokenizer = AutoTokenizer.from_pretrained(self.model_name_or_path)
-
- def _text_length(self, text: Any) -> int:
- """
- Help function to get the length for the input text. Text can be either
- a list of ints (which means a single text as input), or a tuple of list of ints
- (representing several text inputs to the model).
- """
-
- if isinstance(text, dict): # {key: value} case
- return len(next(iter(text.values())))
- elif not hasattr(text, "__len__"): # Object has no len() method
- return 1
- # Empty string or list of ints
- elif len(text) == 0 or isinstance(text[0], int):
- return len(text)
- else:
- # Sum of length of individual strings
- return sum([len(t) for t in text])
-
- def encode(
- self,
- sentences: Any,
- batch_size: int = 4,
- show_progress_bar: bool = False,
- convert_to_numpy: bool = True,
- convert_to_tensor: bool = False,
- mean_pooling: bool = False,
- normalize_embeddings: bool = True,
- ) -> Any:
- """
- Computes sentence embeddings.
-
- :param sentences: the sentences to embed.
- :param batch_size: the batch size used for the computation.
- :param show_progress_bar: Whether to output a progress bar.
- :param convert_to_numpy: Whether the output should be a list of numpy vectors.
- :param convert_to_tensor: Whether the output should be one large tensor.
- :param mean_pooling: Whether to pool returned vectors.
- :param normalize_embeddings: Whether to normalize returned vectors.
-
- :return: By default, a 2d numpy array with shape [num_inputs, output_dimension].
- """
- try:
- import numpy as np
- except ImportError as e:
- raise ImportError(
- "Unable to import numpy, please install with `pip install -U numpy`."
- ) from e
- try:
- from tqdm import trange
- except ImportError as e:
- raise ImportError(
- "Unable to import tqdm, please install with `pip install -U tqdm`."
- ) from e
- try:
- import torch
- except ImportError as e:
- raise ImportError(
- "Unable to import torch, please install with `pip install -U torch`."
- ) from e
-
- def run_mean_pooling(model_output: Any, attention_mask: Any) -> Any:
- token_embeddings = model_output[
- 0
- ] # First element of model_output contains all token embeddings
- input_mask_expanded = (
- attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()
- )
- return torch.sum(token_embeddings * input_mask_expanded, 1) / torch.clamp(
- input_mask_expanded.sum(1), min=1e-9
- )
-
- if convert_to_tensor:
- convert_to_numpy = False
-
- input_was_string = False
- if isinstance(sentences, str) or not hasattr(
- sentences, "__len__"
- ): # Cast an individual sentence to a list with length 1
- sentences = [sentences]
- input_was_string = True
-
- all_embeddings: Any = []
- length_sorted_idx = np.argsort([-self._text_length(sen) for sen in sentences])
- sentences_sorted = [sentences[idx] for idx in length_sorted_idx]
-
- for start_index in trange(
- 0, len(sentences), batch_size, desc="Batches", disable=not show_progress_bar
- ):
- sentences_batch = sentences_sorted[start_index : start_index + batch_size]
-
- length = self.ov_model.request.inputs[0].get_partial_shape()[1]
- if length.is_dynamic:
- features = self.tokenizer(
- sentences_batch, padding=True, truncation=True, return_tensors="pt"
- )
- else:
- features = self.tokenizer(
- sentences_batch,
- padding="max_length",
- max_length=length.get_length(),
- truncation=True,
- return_tensors="pt",
- )
-
- out_features = self.ov_model(**features)
- if mean_pooling:
- embeddings = run_mean_pooling(out_features, features["attention_mask"])
- else:
- embeddings = out_features[0][:, 0]
- if normalize_embeddings:
- embeddings = torch.nn.functional.normalize(embeddings, p=2, dim=1)
-
- # fixes for #522 and #487 to avoid oom problems on gpu with large datasets
- if convert_to_numpy:
- embeddings = embeddings.cpu()
-
- all_embeddings.extend(embeddings)
-
- all_embeddings = [all_embeddings[idx] for idx in np.argsort(length_sorted_idx)]
-
- if convert_to_tensor:
- if len(all_embeddings):
- all_embeddings = torch.stack(all_embeddings)
- else:
- all_embeddings = torch.Tensor()
- elif convert_to_numpy:
- all_embeddings = np.asarray([emb.numpy() for emb in all_embeddings])
-
- if input_was_string:
- all_embeddings = all_embeddings[0]
-
- return all_embeddings
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Compute doc embeddings using a HuggingFace transformer model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
-
- texts = list(map(lambda x: x.replace("\n", " "), texts))
- embeddings = self.encode(
- texts, show_progress_bar=self.show_progress, **self.encode_kwargs
- )
-
- return embeddings.tolist()
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using a HuggingFace transformer model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self.embed_documents([text])[0]
-
- def save_model(
- self,
- model_path: str,
- ) -> bool:
- self.ov_model.half()
- self.ov_model.save_pretrained(model_path)
- self.tokenizer.save_pretrained(model_path)
- return True
-
-
-class OpenVINOBgeEmbeddings(OpenVINOEmbeddings):
- """OpenVNO BGE embedding models.
-
- Bge Example:
- .. code-block:: python
-
- from langchain_community.embeddings import OpenVINOBgeEmbeddings
-
- model_name = "BAAI/bge-large-en-v1.5"
- model_kwargs = {'device': 'CPU'}
- encode_kwargs = {'normalize_embeddings': True}
- ov = OpenVINOBgeEmbeddings(
- model_name_or_path=model_name,
- model_kwargs=model_kwargs,
- encode_kwargs=encode_kwargs
- )
- """
-
- query_instruction: str = DEFAULT_QUERY_BGE_INSTRUCTION_EN
- """Instruction to use for embedding query."""
- embed_instruction: str = ""
- """Instruction to use for embedding document."""
-
- def __init__(self, **kwargs: Any):
- """Initialize the sentence_transformer."""
- super().__init__(**kwargs)
-
- if "-zh" in self.model_name_or_path:
- self.query_instruction = DEFAULT_QUERY_BGE_INSTRUCTION_ZH
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Compute doc embeddings using a HuggingFace transformer model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- texts = [self.embed_instruction + t.replace("\n", " ") for t in texts]
- embeddings = self.encode(texts, **self.encode_kwargs)
- return embeddings.tolist()
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using a HuggingFace transformer model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- text = text.replace("\n", " ")
- embedding = self.encode(self.query_instruction + text, **self.encode_kwargs)
- return embedding.tolist()
diff --git a/libs/community/langchain_community/embeddings/optimum_intel.py b/libs/community/langchain_community/embeddings/optimum_intel.py
deleted file mode 100644
index 05de5cca3a..0000000000
--- a/libs/community/langchain_community/embeddings/optimum_intel.py
+++ /dev/null
@@ -1,208 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict
-
-
-class QuantizedBiEncoderEmbeddings(BaseModel, Embeddings):
- """Quantized bi-encoders embedding models.
-
- Please ensure that you have installed optimum-intel and ipex.
-
- Input:
- model_name: str = Model name.
- max_seq_len: int = The maximum sequence length for tokenization. (default 512)
- pooling_strategy: str =
- "mean" or "cls", pooling strategy for the final layer. (default "mean")
- query_instruction: Optional[str] =
- An instruction to add to the query before embedding. (default None)
- document_instruction: Optional[str] =
- An instruction to add to each document before embedding. (default None)
- padding: Optional[bool] =
- Whether to add padding during tokenization or not. (default True)
- model_kwargs: Optional[Dict] =
- Parameters to add to the model during initialization. (default {})
- encode_kwargs: Optional[Dict] =
- Parameters to add during the embedding forward pass. (default {})
-
- Example:
-
- from langchain_community.embeddings import QuantizedBiEncoderEmbeddings
-
- model_name = "Intel/bge-small-en-v1.5-rag-int8-static"
- encode_kwargs = {'normalize_embeddings': True}
- hf = QuantizedBiEncoderEmbeddings(
- model_name,
- encode_kwargs=encode_kwargs,
- query_instruction="Represent this sentence for searching relevant passages: "
- )
- """
-
- def __init__(
- self,
- model_name: str,
- max_seq_len: int = 512,
- pooling_strategy: str = "mean", # "mean" or "cls"
- query_instruction: Optional[str] = None,
- document_instruction: Optional[str] = None,
- padding: bool = True,
- model_kwargs: Optional[Dict] = None,
- encode_kwargs: Optional[Dict] = None,
- **kwargs: Any,
- ) -> None:
- super().__init__(**kwargs)
- self.model_name_or_path = model_name
- self.max_seq_len = max_seq_len
- self.pooling = pooling_strategy
- self.padding = padding
- self.encode_kwargs = encode_kwargs or {}
- self.model_kwargs = model_kwargs or {}
-
- self.normalize = self.encode_kwargs.get("normalize_embeddings", False)
- self.batch_size = self.encode_kwargs.get("batch_size", 32)
-
- self.query_instruction = query_instruction
- self.document_instruction = document_instruction
-
- self.load_model()
-
- def load_model(self) -> None:
- try:
- from transformers import AutoTokenizer
- except ImportError as e:
- raise ImportError(
- "Unable to import transformers, please install with "
- "`pip install -U transformers`."
- ) from e
- try:
- from optimum.intel import IPEXModel
-
- self.transformer_model = IPEXModel.from_pretrained(
- self.model_name_or_path, **self.model_kwargs
- )
- except Exception as e:
- raise Exception(
- f"""
-Failed to load model {self.model_name_or_path}, due to the following error:
-{e}
-Please ensure that you have installed optimum-intel and ipex correctly,using:
-
-pip install optimum[neural-compressor]
-pip install intel_extension_for_pytorch
-
-For more information, please visit:
-* Install optimum-intel as shown here: https://github.com/huggingface/optimum-intel.
-* Install IPEX as shown here: https://intel.github.io/intel-extension-for-pytorch/index.html#installation?platform=cpu&version=v2.2.0%2Bcpu.
-"""
- )
- self.transformer_tokenizer = AutoTokenizer.from_pretrained(
- pretrained_model_name_or_path=self.model_name_or_path,
- )
- self.transformer_model.eval()
-
- model_config = ConfigDict(
- extra="allow",
- protected_namespaces=(),
- )
-
- def _embed(self, inputs: Any) -> Any:
- try:
- import torch
- except ImportError as e:
- raise ImportError(
- "Unable to import torch, please install with `pip install -U torch`."
- ) from e
- with torch.inference_mode():
- outputs = self.transformer_model(**inputs)
- if self.pooling == "mean":
- emb = self._mean_pooling(outputs, inputs["attention_mask"])
- elif self.pooling == "cls":
- emb = self._cls_pooling(outputs)
- else:
- raise ValueError("pooling method no supported")
-
- if self.normalize:
- emb = torch.nn.functional.normalize(emb, p=2, dim=1)
- return emb
-
- @staticmethod
- def _cls_pooling(outputs: Any) -> Any:
- if isinstance(outputs, dict):
- token_embeddings = outputs["last_hidden_state"]
- else:
- token_embeddings = outputs[0]
- return token_embeddings[:, 0]
-
- @staticmethod
- def _mean_pooling(outputs: Any, attention_mask: Any) -> Any:
- try:
- import torch
- except ImportError as e:
- raise ImportError(
- "Unable to import torch, please install with `pip install -U torch`."
- ) from e
- if isinstance(outputs, dict):
- token_embeddings = outputs["last_hidden_state"]
- else:
- # First element of model_output contains all token embeddings
- token_embeddings = outputs[0]
- input_mask_expanded = (
- attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()
- )
- sum_embeddings = torch.sum(token_embeddings * input_mask_expanded, 1)
- sum_mask = torch.clamp(input_mask_expanded.sum(1), min=1e-9)
- return sum_embeddings / sum_mask
-
- def _embed_text(self, texts: List[str]) -> List[List[float]]:
- inputs = self.transformer_tokenizer(
- texts,
- max_length=self.max_seq_len,
- truncation=True,
- padding=self.padding,
- return_tensors="pt",
- )
- return self._embed(inputs).tolist()
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed a list of text documents using the Optimized Embedder model.
-
- Input:
- texts: List[str] = List of text documents to embed.
- Output:
- List[List[float]] = The embeddings of each text document.
- """
- try:
- import pandas as pd
- except ImportError as e:
- raise ImportError(
- "Unable to import pandas, please install with `pip install -U pandas`."
- ) from e
- try:
- from tqdm import tqdm
- except ImportError as e:
- raise ImportError(
- "Unable to import tqdm, please install with `pip install -U tqdm`."
- ) from e
- docs = [
- self.document_instruction + d if self.document_instruction else d
- for d in texts
- ]
-
- # group into batches
- text_list_df = pd.DataFrame(docs, columns=["texts"]).reset_index()
-
- # assign each example with its batch
- text_list_df["batch_index"] = text_list_df["index"] // self.batch_size
-
- # create groups
- batches = list(text_list_df.groupby(["batch_index"])["texts"].apply(list))
-
- vectors = []
- for batch in tqdm(batches, desc="Batches"):
- vectors += self._embed_text(batch)
- return vectors
-
- def embed_query(self, text: str) -> List[float]:
- if self.query_instruction:
- text = self.query_instruction + text
- return self._embed_text([text])[0]
diff --git a/libs/community/langchain_community/embeddings/oracleai.py b/libs/community/langchain_community/embeddings/oracleai.py
deleted file mode 100644
index d1dca41905..0000000000
--- a/libs/community/langchain_community/embeddings/oracleai.py
+++ /dev/null
@@ -1,194 +0,0 @@
-# Authors:
-# Harichandan Roy (hroy)
-# David Jiang (ddjiang)
-#
-# -----------------------------------------------------------------------------
-# oracleai.py
-# -----------------------------------------------------------------------------
-
-from __future__ import annotations
-
-import json
-import logging
-import traceback
-from typing import TYPE_CHECKING, Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict
-
-if TYPE_CHECKING:
- from oracledb import Connection
-
-logger = logging.getLogger(__name__)
-
-"""OracleEmbeddings class"""
-
-
-class OracleEmbeddings(BaseModel, Embeddings):
- """Get Embeddings"""
-
- """Oracle Connection"""
- conn: Any = None
- """Embedding Parameters"""
- params: Dict[str, Any]
- """Proxy"""
- proxy: Optional[str] = None
-
- def __init__(self, **kwargs: Any):
- super().__init__(**kwargs)
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- """
- 1 - user needs to have create procedure,
- create mining model, create any directory privilege.
- 2 - grant create procedure, create mining model,
- create any directory to ;
- """
-
- @staticmethod
- def load_onnx_model(
- conn: Connection, dir: str, onnx_file: str, model_name: str
- ) -> None:
- """Load an ONNX model to Oracle Database.
- Args:
- conn: Oracle Connection,
- dir: Oracle Directory,
- onnx_file: ONNX file name,
- model_name: Name of the model.
- """
-
- try:
- if conn is None or dir is None or onnx_file is None or model_name is None:
- raise Exception("Invalid input")
-
- cursor = conn.cursor()
- cursor.execute(
- """
- begin
- dbms_data_mining.drop_model(model_name => :model, force => true);
- SYS.DBMS_VECTOR.load_onnx_model(:path, :filename, :model,
- json('{"function" : "embedding",
- "embeddingOutput" : "embedding",
- "input": {"input": ["DATA"]}}'));
- end;""",
- path=dir,
- filename=onnx_file,
- model=model_name,
- )
-
- cursor.close()
-
- except Exception as ex:
- logger.info(f"An exception occurred :: {ex}")
- traceback.print_exc()
- cursor.close()
- raise
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Compute doc embeddings using an OracleEmbeddings.
- Args:
- texts: The list of texts to embed.
- Returns:
- List of embeddings, one for each input text.
- """
-
- try:
- import oracledb
- except ImportError as e:
- raise ImportError(
- "Unable to import oracledb, please install with "
- "`pip install -U oracledb`."
- ) from e
-
- if texts is None:
- return None
-
- embeddings: List[List[float]] = []
- try:
- # returns strings or bytes instead of a locator
- oracledb.defaults.fetch_lobs = False
- cursor = self.conn.cursor()
-
- if self.proxy:
- cursor.execute(
- "begin utl_http.set_proxy(:proxy); end;", proxy=self.proxy
- )
-
- chunks = []
- for i, text in enumerate(texts, start=1):
- chunk = {"chunk_id": i, "chunk_data": text}
- chunks.append(json.dumps(chunk))
-
- vector_array_type = self.conn.gettype("SYS.VECTOR_ARRAY_T")
- inputs = vector_array_type.newobject(chunks)
- cursor.execute(
- "select t.* "
- + "from dbms_vector_chain.utl_to_embeddings(:content, "
- + "json(:params)) t",
- content=inputs,
- params=json.dumps(self.params),
- )
-
- for row in cursor:
- if row is None:
- embeddings.append([])
- else:
- rdata = json.loads(row[0])
- # dereference string as array
- vec = json.loads(rdata["embed_vector"])
- embeddings.append(vec)
-
- cursor.close()
- return embeddings
- except Exception as ex:
- logger.info(f"An exception occurred :: {ex}")
- traceback.print_exc()
- cursor.close()
- raise
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embedding using an OracleEmbeddings.
- Args:
- text: The text to embed.
- Returns:
- Embedding for the text.
- """
- return self.embed_documents([text])[0]
-
-
-# uncomment the following code block to run the test
-
-"""
-# A sample unit test.
-
-import oracledb
-# get the Oracle connection
-conn = oracledb.connect(
- user="",
- password="",
- dsn="/",
-)
-print("Oracle connection is established...")
-
-# params
-embedder_params = {"provider": "database", "model": "demo_model"}
-proxy = ""
-
-# instance
-embedder = OracleEmbeddings(conn=conn, params=embedder_params, proxy=proxy)
-
-docs = ["hello world!", "hi everyone!", "greetings!"]
-embeds = embedder.embed_documents(docs)
-print(f"Total Embeddings: {len(embeds)}")
-print(f"Embedding generated by OracleEmbeddings: {embeds[0]}\n")
-
-embed = embedder.embed_query("Hello World!")
-print(f"Embedding generated by OracleEmbeddings: {embed}")
-
-conn.close()
-print("Connection is closed.")
-
-"""
diff --git a/libs/community/langchain_community/embeddings/ovhcloud.py b/libs/community/langchain_community/embeddings/ovhcloud.py
deleted file mode 100644
index 49b9cbfa21..0000000000
--- a/libs/community/langchain_community/embeddings/ovhcloud.py
+++ /dev/null
@@ -1,115 +0,0 @@
-import json
-import logging
-import time
-from typing import Any, List
-
-import requests
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict
-
-logger = logging.getLogger(__name__)
-
-
-class OVHCloudEmbeddings(BaseModel, Embeddings):
- """
- OVHcloud AI Endpoints Embeddings.
- """
-
- """ OVHcloud AI Endpoints Access Token"""
- access_token: str = ""
-
- """ OVHcloud AI Endpoints model name for embeddings generation"""
- model_name: str = ""
-
- """ OVHcloud AI Endpoints region"""
- region: str = "kepler"
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- def __init__(self, **kwargs: Any):
- super().__init__(**kwargs)
- if self.access_token == "":
- raise ValueError("Access token is required for OVHCloud embeddings.")
- if self.model_name == "":
- raise ValueError("Model name is required for OVHCloud embeddings.")
- if self.region == "":
- raise ValueError("Region is required for OVHCloud embeddings.")
-
- def _generate_embedding(self, text: str) -> List[float]:
- """Generate embeddings from OVHCLOUD AIE.
- Args:
- text (str): The text to embed.
- Returns:
- List[float]: Embeddings for the text.
- """
-
- return self._send_request_to_ai_endpoints("text/plain", text, "text2vec")
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed a list of documents.
- Args:
- texts (List[str]): The list of texts to embed.
-
- Returns:
- List[List[float]]: List of embeddings, one for each input text.
-
- """
-
- return self._send_request_to_ai_endpoints(
- "application/json", json.dumps(texts), "batch_text2vec"
- )
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a single query text.
- Args:
- text (str): The text to embed.
- Returns:
- List[float]: Embeddings for the text.
- """
- return self._generate_embedding(text)
-
- def _send_request_to_ai_endpoints(
- self, contentType: str, payload: str, route: str
- ) -> Any:
- """Send a HTTPS request to OVHcloud AI Endpoints
- Args:
- contentType (str): The content type of the request, application/json or text/plain.
- payload (str): The payload of the request.
- route (str): The route of the request, batch_text2vec or text2vec.
- """ # noqa: E501
- headers = {
- "content-type": contentType,
- "Authorization": f"Bearer {self.access_token}",
- }
-
- session = requests.session()
- while True:
- response = session.post(
- (
- f"https://{self.model_name}.endpoints.{self.region}"
- f".ai.cloud.ovh.net/api/{route}"
- ),
- headers=headers,
- data=payload,
- )
- if response.status_code != 200:
- if response.status_code == 429:
- """Rate limit exceeded, wait for reset"""
- reset_time = int(response.headers.get("RateLimit-Reset", 0))
- logger.info("Rate limit exceeded. Waiting %d seconds.", reset_time)
- if reset_time > 0:
- time.sleep(reset_time)
- continue
- else:
- """Rate limit reset time has passed, retry immediately"""
- continue
- if response.status_code == 401:
- """ Unauthorized, retry with new token """
- raise ValueError("Unauthorized, retry with new token")
- """ Handle other non-200 status codes """
- raise ValueError(
- "Request failed with status code: {status_code}, {text}".format(
- status_code=response.status_code, text=response.text
- )
- )
- return response.json()
diff --git a/libs/community/langchain_community/embeddings/premai.py b/libs/community/langchain_community/embeddings/premai.py
deleted file mode 100644
index ed4a745344..0000000000
--- a/libs/community/langchain_community/embeddings/premai.py
+++ /dev/null
@@ -1,130 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Callable, Dict, List, Optional, Union
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.language_models.llms import create_base_retry_decorator
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import BaseModel, SecretStr
-
-logger = logging.getLogger(__name__)
-
-
-class PremAIEmbeddings(BaseModel, Embeddings):
- """Prem's Embedding APIs"""
-
- project_id: int
- """The project ID in which the experiments or deployments are carried out.
- You can find all your projects here: https://app.premai.io/projects/"""
-
- premai_api_key: Optional[SecretStr] = None
- """Prem AI API Key. Get it here: https://app.premai.io/api_keys/"""
-
- model: str
- """The Embedding model to choose from"""
-
- show_progress_bar: bool = False
- """Whether to show a tqdm progress bar. Must have `tqdm` installed."""
-
- max_retries: int = 1
- """Max number of retries for tenacity"""
-
- client: Any
-
- @pre_init
- def validate_environments(cls, values: Dict) -> Dict:
- """Validate that the package is installed and that the API token is valid"""
- try:
- from premai import Prem
- except ImportError as error:
- raise ImportError(
- "Could not import Prem Python package."
- "Please install it with: `pip install premai`"
- ) from error
-
- try:
- premai_api_key = get_from_dict_or_env(
- values, "premai_api_key", "PREMAI_API_KEY"
- )
- values["client"] = Prem(api_key=premai_api_key)
- except Exception as error:
- raise ValueError("Your API Key is incorrect. Please try again.") from error
- return values
-
- def embed_query(self, text: str) -> List[float]:
- """Embed query text"""
- embeddings = embed_with_retry(
- self, model=self.model, project_id=self.project_id, input=text
- )
- return embeddings.data[0].embedding
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- embeddings = embed_with_retry(
- self, model=self.model, project_id=self.project_id, input=texts
- ).data
-
- return [embedding.embedding for embedding in embeddings]
-
-
-def create_prem_retry_decorator(
- embedder: PremAIEmbeddings,
- *,
- max_retries: int = 1,
-) -> Callable[[Any], Any]:
- """Create a retry decorator for PremAIEmbeddings.
-
- Args:
- embedder (PremAIEmbeddings): The PremAIEmbeddings instance
- max_retries (int): The maximum number of retries
-
- Returns:
- Callable[[Any], Any]: The retry decorator
- """
- import premai.models
-
- errors = [
- premai.models.api_response_validation_error.APIResponseValidationError,
- premai.models.conflict_error.ConflictError,
- premai.models.model_not_found_error.ModelNotFoundError,
- premai.models.permission_denied_error.PermissionDeniedError,
- premai.models.provider_api_connection_error.ProviderAPIConnectionError,
- premai.models.provider_api_status_error.ProviderAPIStatusError,
- premai.models.provider_api_timeout_error.ProviderAPITimeoutError,
- premai.models.provider_internal_server_error.ProviderInternalServerError,
- premai.models.provider_not_found_error.ProviderNotFoundError,
- premai.models.rate_limit_error.RateLimitError,
- premai.models.unprocessable_entity_error.UnprocessableEntityError,
- premai.models.validation_error.ValidationError,
- ]
-
- decorator = create_base_retry_decorator(
- error_types=errors, max_retries=max_retries, run_manager=None
- )
- return decorator
-
-
-def embed_with_retry(
- embedder: PremAIEmbeddings,
- model: str,
- project_id: int,
- input: Union[str, List[str]],
-) -> Any:
- """Using tenacity for retry in embedding calls"""
- retry_decorator = create_prem_retry_decorator(
- embedder, max_retries=embedder.max_retries
- )
-
- @retry_decorator
- def _embed_with_retry(
- embedder: PremAIEmbeddings,
- project_id: int,
- model: str,
- input: Union[str, List[str]],
- ) -> Any:
- embedding_response = embedder.client.embeddings.create(
- project_id=project_id, model=model, input=input
- )
- return embedding_response
-
- return _embed_with_retry(embedder, project_id=project_id, model=model, input=input)
diff --git a/libs/community/langchain_community/embeddings/sagemaker_endpoint.py b/libs/community/langchain_community/embeddings/sagemaker_endpoint.py
deleted file mode 100644
index d69cbd92ea..0000000000
--- a/libs/community/langchain_community/embeddings/sagemaker_endpoint.py
+++ /dev/null
@@ -1,210 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import pre_init
-from pydantic import BaseModel, ConfigDict
-
-from langchain_community.llms.sagemaker_endpoint import ContentHandlerBase
-
-
-class EmbeddingsContentHandler(ContentHandlerBase[List[str], List[List[float]]]):
- """Content handler for LLM class."""
-
-
-class SagemakerEndpointEmbeddings(BaseModel, Embeddings):
- """Custom Sagemaker Inference Endpoints.
-
- To use, you must supply the endpoint name from your deployed
- Sagemaker model & the region where it is deployed.
-
- To authenticate, the AWS client uses the following methods to
- automatically load credentials:
- https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html
-
- If a specific credential profile should be used, you must pass
- the name of the profile from the ~/.aws/credentials file that is to be used.
-
- Make sure the credentials / roles used have the required policies to
- access the Sagemaker endpoint.
- See: https://docs.aws.amazon.com/IAM/latest/UserGuide/access_policies.html
- """
-
- """
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import SagemakerEndpointEmbeddings
- endpoint_name = (
- "my-endpoint-name"
- )
- region_name = (
- "us-west-2"
- )
- credentials_profile_name = (
- "default"
- )
- se = SagemakerEndpointEmbeddings(
- endpoint_name=endpoint_name,
- region_name=region_name,
- credentials_profile_name=credentials_profile_name
- )
-
- #Use with boto3 client
- client = boto3.client(
- "sagemaker-runtime",
- region_name=region_name
- )
- se = SagemakerEndpointEmbeddings(
- endpoint_name=endpoint_name,
- client=client
- )
- """
- client: Any = None
-
- endpoint_name: str = ""
- """The name of the endpoint from the deployed Sagemaker model.
- Must be unique within an AWS Region."""
-
- region_name: str = ""
- """The aws region where the Sagemaker model is deployed, eg. `us-west-2`."""
-
- credentials_profile_name: Optional[str] = None
- """The name of the profile in the ~/.aws/credentials or ~/.aws/config files, which
- has either access keys or role information specified.
- If not specified, the default credential profile or, if on an EC2 instance,
- credentials from IMDS will be used.
- See: https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html
- """
-
- content_handler: EmbeddingsContentHandler
- """The content handler class that provides an input and
- output transform functions to handle formats between LLM
- and the endpoint.
- """
-
- """
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings.sagemaker_endpoint import EmbeddingsContentHandler
-
- class ContentHandler(EmbeddingsContentHandler):
- content_type = "application/json"
- accepts = "application/json"
-
- def transform_input(self, prompts: List[str], model_kwargs: Dict) -> bytes:
- input_str = json.dumps({prompts: prompts, **model_kwargs})
- return input_str.encode('utf-8')
-
- def transform_output(self, output: bytes) -> List[List[float]]:
- response_json = json.loads(output.read().decode("utf-8"))
- return response_json["vectors"]
- """ # noqa: E501
-
- model_kwargs: Optional[Dict] = None
- """Keyword arguments to pass to the model."""
-
- endpoint_kwargs: Optional[Dict] = None
- """Optional attributes passed to the invoke_endpoint
- function. See `boto3`_. docs for more info.
- .. _boto3:
- """
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True, extra="forbid", protected_namespaces=()
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Dont do anything if client provided externally"""
- if values.get("client") is not None:
- return values
-
- """Validate that AWS credentials to and python package exists in environment."""
- try:
- import boto3
-
- try:
- if values["credentials_profile_name"] is not None:
- session = boto3.Session(
- profile_name=values["credentials_profile_name"]
- )
- else:
- # use default credentials
- session = boto3.Session()
-
- values["client"] = session.client(
- "sagemaker-runtime", region_name=values["region_name"]
- )
-
- except Exception as e:
- raise ValueError(
- "Could not load credentials to authenticate with AWS client. "
- "Please check that credentials in the specified "
- f"profile name are valid. {e}"
- ) from e
-
- except ImportError:
- raise ImportError(
- "Could not import boto3 python package. "
- "Please install it with `pip install boto3`."
- )
- return values
-
- def _embedding_func(self, texts: List[str]) -> List[List[float]]:
- """Call out to SageMaker Inference embedding endpoint."""
- # replace newlines, which can negatively affect performance.
- texts = list(map(lambda x: x.replace("\n", " "), texts))
- _model_kwargs = self.model_kwargs or {}
- _endpoint_kwargs = self.endpoint_kwargs or {}
-
- body = self.content_handler.transform_input(texts, _model_kwargs)
- content_type = self.content_handler.content_type
- accepts = self.content_handler.accepts
-
- # send request
- try:
- response = self.client.invoke_endpoint(
- EndpointName=self.endpoint_name,
- Body=body,
- ContentType=content_type,
- Accept=accepts,
- **_endpoint_kwargs,
- )
- except Exception as e:
- raise ValueError(f"Error raised by inference endpoint: {e}")
-
- return self.content_handler.transform_output(response["Body"])
-
- def embed_documents(
- self, texts: List[str], chunk_size: int = 64
- ) -> List[List[float]]:
- """Compute doc embeddings using a SageMaker Inference Endpoint.
-
- Args:
- texts: The list of texts to embed.
- chunk_size: The chunk size defines how many input texts will
- be grouped together as request. If None, will use the
- chunk size specified by the class.
-
-
- Returns:
- List of embeddings, one for each text.
- """
- results = []
- _chunk_size = len(texts) if chunk_size > len(texts) else chunk_size
- for i in range(0, len(texts), _chunk_size):
- response = self._embedding_func(texts[i : i + _chunk_size])
- results.extend(response)
- return results
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using a SageMaker inference endpoint.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return self._embedding_func([text])[0]
diff --git a/libs/community/langchain_community/embeddings/sambanova.py b/libs/community/langchain_community/embeddings/sambanova.py
deleted file mode 100644
index e9d76cae87..0000000000
--- a/libs/community/langchain_community/embeddings/sambanova.py
+++ /dev/null
@@ -1,324 +0,0 @@
-import json
-from typing import Dict, Generator, List, Optional
-
-import requests
-from langchain_core._api.deprecation import deprecated
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict
-
-
-@deprecated(
- since="0.3.16",
- removal="1.0",
- alternative_import="langchain_sambanova.SambaStudioEmbeddings",
-)
-class SambaStudioEmbeddings(BaseModel, Embeddings):
- """SambaNova embedding models.
-
- To use, you should have the environment variables
- ``SAMBASTUDIO_EMBEDDINGS_BASE_URL``, ``SAMBASTUDIO_EMBEDDINGS_BASE_URI``
- ``SAMBASTUDIO_EMBEDDINGS_PROJECT_ID``, ``SAMBASTUDIO_EMBEDDINGS_ENDPOINT_ID``,
- ``SAMBASTUDIO_EMBEDDINGS_API_KEY``
- set with your personal sambastudio variable or pass it as a named parameter
- to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import SambaStudioEmbeddings
-
- embeddings = SambaStudioEmbeddings(sambastudio_embeddings_base_url=base_url,
- sambastudio_embeddings_base_uri=base_uri,
- sambastudio_embeddings_project_id=project_id,
- sambastudio_embeddings_endpoint_id=endpoint_id,
- sambastudio_embeddings_api_key=api_key,
- batch_size=32)
- (or)
-
- embeddings = SambaStudioEmbeddings(batch_size=32)
-
- (or)
-
- # CoE example
- embeddings = SambaStudioEmbeddings(
- batch_size=1,
- model_kwargs={
- 'select_expert':'e5-mistral-7b-instruct'
- }
- )
- """
-
- sambastudio_embeddings_base_url: str = ""
- """Base url to use"""
-
- sambastudio_embeddings_base_uri: str = ""
- """endpoint base uri"""
-
- sambastudio_embeddings_project_id: str = ""
- """Project id on sambastudio for model"""
-
- sambastudio_embeddings_endpoint_id: str = ""
- """endpoint id on sambastudio for model"""
-
- sambastudio_embeddings_api_key: str = ""
- """sambastudio api key"""
-
- model_kwargs: dict = {}
- """Key word arguments to pass to the model."""
-
- batch_size: int = 32
- """Batch size for the embedding models"""
-
- model_config = ConfigDict(protected_namespaces=())
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- values["sambastudio_embeddings_base_url"] = get_from_dict_or_env(
- values, "sambastudio_embeddings_base_url", "SAMBASTUDIO_EMBEDDINGS_BASE_URL"
- )
- values["sambastudio_embeddings_base_uri"] = get_from_dict_or_env(
- values,
- "sambastudio_embeddings_base_uri",
- "SAMBASTUDIO_EMBEDDINGS_BASE_URI",
- default="api/predict/generic",
- )
- values["sambastudio_embeddings_project_id"] = get_from_dict_or_env(
- values,
- "sambastudio_embeddings_project_id",
- "SAMBASTUDIO_EMBEDDINGS_PROJECT_ID",
- )
- values["sambastudio_embeddings_endpoint_id"] = get_from_dict_or_env(
- values,
- "sambastudio_embeddings_endpoint_id",
- "SAMBASTUDIO_EMBEDDINGS_ENDPOINT_ID",
- )
- values["sambastudio_embeddings_api_key"] = get_from_dict_or_env(
- values, "sambastudio_embeddings_api_key", "SAMBASTUDIO_EMBEDDINGS_API_KEY"
- )
- return values
-
- def _get_tuning_params(self) -> str:
- """
- Get the tuning parameters to use when calling the model
-
- Returns:
- The tuning parameters as a JSON string.
- """
- if "api/v2/predict/generic" in self.sambastudio_embeddings_base_uri:
- tuning_params_dict = self.model_kwargs
- else:
- tuning_params_dict = {
- k: {"type": type(v).__name__, "value": str(v)}
- for k, v in (self.model_kwargs.items())
- }
- tuning_params = json.dumps(tuning_params_dict)
- return tuning_params
-
- def _get_full_url(self, path: str) -> str:
- """
- Return the full API URL for a given path.
-
- :param str path: the sub-path
- :returns: the full API URL for the sub-path
- :rtype: str
- """
- return f"{self.sambastudio_embeddings_base_url}/{self.sambastudio_embeddings_base_uri}/{path}" # noqa: E501
-
- def _iterate_over_batches(self, texts: List[str], batch_size: int) -> Generator:
- """Generator for creating batches in the embed documents method
- Args:
- texts (List[str]): list of strings to embed
- batch_size (int, optional): batch size to be used for the embedding model.
- Will depend on the RDU endpoint used.
- Yields:
- List[str]: list (batch) of strings of size batch size
- """
- for i in range(0, len(texts), batch_size):
- yield texts[i : i + batch_size]
-
- def embed_documents(
- self, texts: List[str], batch_size: Optional[int] = None
- ) -> List[List[float]]:
- """Returns a list of embeddings for the given sentences.
- Args:
- texts (`List[str]`): List of texts to encode
- batch_size (`int`): Batch size for the encoding
-
- Returns:
- `List[np.ndarray]` or `List[tensor]`: List of embeddings
- for the given sentences
- """
- if batch_size is None:
- batch_size = self.batch_size
- http_session = requests.Session()
- url = self._get_full_url(
- f"{self.sambastudio_embeddings_project_id}/{self.sambastudio_embeddings_endpoint_id}"
- )
- params = json.loads(self._get_tuning_params())
- embeddings = []
-
- if "api/predict/nlp" in self.sambastudio_embeddings_base_uri:
- for batch in self._iterate_over_batches(texts, batch_size):
- data = {"inputs": batch, "params": params}
- response = http_session.post(
- url,
- headers={"key": self.sambastudio_embeddings_api_key},
- json=data,
- )
- if response.status_code != 200:
- raise RuntimeError(
- f"Sambanova /complete call failed with status code "
- f"{response.status_code}.\n Details: {response.text}"
- )
- try:
- embedding = response.json()["data"]
- embeddings.extend(embedding)
- except KeyError:
- raise KeyError(
- "'data' not found in endpoint response",
- response.json(),
- )
-
- elif "api/v2/predict/generic" in self.sambastudio_embeddings_base_uri:
- for batch in self._iterate_over_batches(texts, batch_size):
- items = [
- {"id": f"item{i}", "value": item} for i, item in enumerate(batch)
- ]
- data = {"items": items, "params": params}
- response = http_session.post(
- url,
- headers={"key": self.sambastudio_embeddings_api_key},
- json=data,
- )
- if response.status_code != 200:
- raise RuntimeError(
- f"Sambanova /complete call failed with status code "
- f"{response.status_code}.\n Details: {response.text}"
- )
- try:
- embedding = [item["value"] for item in response.json()["items"]]
- embeddings.extend(embedding)
- except KeyError:
- raise KeyError(
- "'items' not found in endpoint response",
- response.json(),
- )
-
- elif "api/predict/generic" in self.sambastudio_embeddings_base_uri:
- for batch in self._iterate_over_batches(texts, batch_size):
- data = {"instances": batch, "params": params}
- response = http_session.post(
- url,
- headers={"key": self.sambastudio_embeddings_api_key},
- json=data,
- )
- if response.status_code != 200:
- raise RuntimeError(
- f"Sambanova /complete call failed with status code "
- f"{response.status_code}.\n Details: {response.text}"
- )
- try:
- if params.get("select_expert"):
- embedding = response.json()["predictions"]
- else:
- embedding = response.json()["predictions"]
- embeddings.extend(embedding)
- except KeyError:
- raise KeyError(
- "'predictions' not found in endpoint response",
- response.json(),
- )
-
- else:
- raise ValueError(
- f"handling of endpoint uri: {self.sambastudio_embeddings_base_uri} not implemented" # noqa: E501
- )
-
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Returns a list of embeddings for the given sentences.
- Args:
- sentences (`List[str]`): List of sentences to encode
-
- Returns:
- `List[np.ndarray]` or `List[tensor]`: List of embeddings
- for the given sentences
- """
- http_session = requests.Session()
- url = self._get_full_url(
- f"{self.sambastudio_embeddings_project_id}/{self.sambastudio_embeddings_endpoint_id}"
- )
- params = json.loads(self._get_tuning_params())
-
- if "api/predict/nlp" in self.sambastudio_embeddings_base_uri:
- data = {"inputs": [text], "params": params}
- response = http_session.post(
- url,
- headers={"key": self.sambastudio_embeddings_api_key},
- json=data,
- )
- if response.status_code != 200:
- raise RuntimeError(
- f"Sambanova /complete call failed with status code "
- f"{response.status_code}.\n Details: {response.text}"
- )
- try:
- embedding = response.json()["data"][0]
- except KeyError:
- raise KeyError(
- "'data' not found in endpoint response",
- response.json(),
- )
-
- elif "api/v2/predict/generic" in self.sambastudio_embeddings_base_uri:
- data = {"items": [{"id": "item0", "value": text}], "params": params}
- response = http_session.post(
- url,
- headers={"key": self.sambastudio_embeddings_api_key},
- json=data,
- )
- if response.status_code != 200:
- raise RuntimeError(
- f"Sambanova /complete call failed with status code "
- f"{response.status_code}.\n Details: {response.text}"
- )
- try:
- embedding = response.json()["items"][0]["value"]
- except KeyError:
- raise KeyError(
- "'items' not found in endpoint response",
- response.json(),
- )
-
- elif "api/predict/generic" in self.sambastudio_embeddings_base_uri:
- data = {"instances": [text], "params": params}
- response = http_session.post(
- url,
- headers={"key": self.sambastudio_embeddings_api_key},
- json=data,
- )
- if response.status_code != 200:
- raise RuntimeError(
- f"Sambanova /complete call failed with status code "
- f"{response.status_code}.\n Details: {response.text}"
- )
- try:
- if params.get("select_expert"):
- embedding = response.json()["predictions"][0]
- else:
- embedding = response.json()["predictions"][0]
- except KeyError:
- raise KeyError(
- "'predictions' not found in endpoint response",
- response.json(),
- )
-
- else:
- raise ValueError(
- f"handling of endpoint uri: {self.sambastudio_embeddings_base_uri} not implemented" # noqa: E501
- )
-
- return embedding
diff --git a/libs/community/langchain_community/embeddings/self_hosted.py b/libs/community/langchain_community/embeddings/self_hosted.py
deleted file mode 100644
index 8099c7018e..0000000000
--- a/libs/community/langchain_community/embeddings/self_hosted.py
+++ /dev/null
@@ -1,101 +0,0 @@
-from typing import Any, Callable, List
-
-from langchain_core.embeddings import Embeddings
-from pydantic import ConfigDict
-
-from langchain_community.llms.self_hosted import SelfHostedPipeline
-
-
-def _embed_documents(pipeline: Any, *args: Any, **kwargs: Any) -> List[List[float]]:
- """Inference function to send to the remote hardware.
-
- Accepts a sentence_transformer model_id and
- returns a list of embeddings for each document in the batch.
- """
- return pipeline(*args, **kwargs)
-
-
-class SelfHostedEmbeddings(SelfHostedPipeline, Embeddings):
- """Custom embedding models on self-hosted remote hardware.
-
- Supported hardware includes auto-launched instances on AWS, GCP, Azure,
- and Lambda, as well as servers specified
- by IP address and SSH credentials (such as on-prem, or another
- cloud like Paperspace, Coreweave, etc.).
-
- To use, you should have the ``runhouse`` python package installed.
-
- Example using a model load function:
- .. code-block:: python
-
- from langchain_community.embeddings import SelfHostedEmbeddings
- from transformers import AutoModelForCausalLM, AutoTokenizer, pipeline
- import runhouse as rh
-
- gpu = rh.cluster(name="rh-a10x", instance_type="A100:1")
- def get_pipeline():
- model_id = "facebook/bart-large"
- tokenizer = AutoTokenizer.from_pretrained(model_id)
- model = AutoModelForCausalLM.from_pretrained(model_id)
- return pipeline("feature-extraction", model=model, tokenizer=tokenizer)
- embeddings = SelfHostedEmbeddings(
- model_load_fn=get_pipeline,
- hardware=gpu
- model_reqs=["./", "torch", "transformers"],
- )
- Example passing in a pipeline path:
- .. code-block:: python
-
- from langchain_community.embeddings import SelfHostedHFEmbeddings
- import runhouse as rh
- from transformers import pipeline
-
- gpu = rh.cluster(name="rh-a10x", instance_type="A100:1")
- pipeline = pipeline(model="bert-base-uncased", task="feature-extraction")
- rh.blob(pickle.dumps(pipeline),
- path="models/pipeline.pkl").save().to(gpu, path="models")
- embeddings = SelfHostedHFEmbeddings.from_pipeline(
- pipeline="models/pipeline.pkl",
- hardware=gpu,
- model_reqs=["./", "torch", "transformers"],
- )
- """
-
- inference_fn: Callable = _embed_documents
- """Inference function to extract the embeddings on the remote hardware."""
- inference_kwargs: Any = None
- """Any kwargs to pass to the model's inference function."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Compute doc embeddings using a HuggingFace transformer model.
-
- Args:
- texts: The list of texts to embed.s
-
- Returns:
- List of embeddings, one for each text.
- """
- texts = list(map(lambda x: x.replace("\n", " "), texts))
- embeddings = self.client(self.pipeline_ref, texts)
- if not isinstance(embeddings, list):
- return embeddings.tolist()
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using a HuggingFace transformer model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- text = text.replace("\n", " ")
- embeddings = self.client(self.pipeline_ref, text)
- if not isinstance(embeddings, list):
- return embeddings.tolist()
- return embeddings
diff --git a/libs/community/langchain_community/embeddings/self_hosted_hugging_face.py b/libs/community/langchain_community/embeddings/self_hosted_hugging_face.py
deleted file mode 100644
index d45a802492..0000000000
--- a/libs/community/langchain_community/embeddings/self_hosted_hugging_face.py
+++ /dev/null
@@ -1,168 +0,0 @@
-import importlib
-import logging
-from typing import Any, Callable, List, Optional
-
-from langchain_community.embeddings.self_hosted import SelfHostedEmbeddings
-
-DEFAULT_MODEL_NAME = "sentence-transformers/all-mpnet-base-v2"
-DEFAULT_INSTRUCT_MODEL = "hkunlp/instructor-large"
-DEFAULT_EMBED_INSTRUCTION = "Represent the document for retrieval: "
-DEFAULT_QUERY_INSTRUCTION = (
- "Represent the question for retrieving supporting documents: "
-)
-
-logger = logging.getLogger(__name__)
-
-
-def _embed_documents(client: Any, *args: Any, **kwargs: Any) -> List[List[float]]:
- """Inference function to send to the remote hardware.
-
- Accepts a sentence_transformer model_id and
- returns a list of embeddings for each document in the batch.
- """
- return client.encode(*args, **kwargs)
-
-
-def load_embedding_model(model_id: str, instruct: bool = False, device: int = 0) -> Any:
- """Load the embedding model."""
- if not instruct:
- import sentence_transformers
-
- client = sentence_transformers.SentenceTransformer(model_id)
- else:
- from InstructorEmbedding import INSTRUCTOR
-
- client = INSTRUCTOR(model_id)
-
- if importlib.util.find_spec("torch") is not None:
- import torch
-
- cuda_device_count = torch.cuda.device_count()
- if device < -1 or (device >= cuda_device_count):
- raise ValueError(
- f"Got device=={device}, "
- f"device is required to be within [-1, {cuda_device_count})"
- )
- if device < 0 and cuda_device_count > 0:
- logger.warning(
- "Device has %d GPUs available. "
- "Provide device={deviceId} to `from_model_id` to use available"
- "GPUs for execution. deviceId is -1 for CPU and "
- "can be a positive integer associated with CUDA device id.",
- cuda_device_count,
- )
-
- client = client.to(device)
- return client
-
-
-class SelfHostedHuggingFaceEmbeddings(SelfHostedEmbeddings):
- """HuggingFace embedding models on self-hosted remote hardware.
-
- Supported hardware includes auto-launched instances on AWS, GCP, Azure,
- and Lambda, as well as servers specified
- by IP address and SSH credentials (such as on-prem, or another cloud
- like Paperspace, Coreweave, etc.).
-
- To use, you should have the ``runhouse`` python package installed.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import SelfHostedHuggingFaceEmbeddings
- import runhouse as rh
- model_id = "sentence-transformers/all-mpnet-base-v2"
- gpu = rh.cluster(name="rh-a10x", instance_type="A100:1")
- hf = SelfHostedHuggingFaceEmbeddings(model_id=model_id, hardware=gpu)
- """
-
- client: Any #: :meta private:
- model_id: str = DEFAULT_MODEL_NAME
- """Model name to use."""
- model_reqs: List[str] = ["./", "sentence_transformers", "torch"]
- """Requirements to install on hardware to inference the model."""
- hardware: Any
- """Remote hardware to send the inference function to."""
- model_load_fn: Callable = load_embedding_model
- """Function to load the model remotely on the server."""
- load_fn_kwargs: Optional[dict] = None
- """Keyword arguments to pass to the model load function."""
- inference_fn: Callable = _embed_documents
- """Inference function to extract the embeddings."""
-
- def __init__(self, **kwargs: Any):
- """Initialize the remote inference function."""
- load_fn_kwargs = kwargs.pop("load_fn_kwargs", {})
- load_fn_kwargs["model_id"] = load_fn_kwargs.get("model_id", DEFAULT_MODEL_NAME)
- load_fn_kwargs["instruct"] = load_fn_kwargs.get("instruct", False)
- load_fn_kwargs["device"] = load_fn_kwargs.get("device", 0)
- super().__init__(load_fn_kwargs=load_fn_kwargs, **kwargs)
-
-
-class SelfHostedHuggingFaceInstructEmbeddings(SelfHostedHuggingFaceEmbeddings):
- """HuggingFace InstructEmbedding models on self-hosted remote hardware.
-
- Supported hardware includes auto-launched instances on AWS, GCP, Azure,
- and Lambda, as well as servers specified
- by IP address and SSH credentials (such as on-prem, or another
- cloud like Paperspace, Coreweave, etc.).
-
- To use, you should have the ``runhouse`` python package installed.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import SelfHostedHuggingFaceInstructEmbeddings
- import runhouse as rh
- model_name = "hkunlp/instructor-large"
- gpu = rh.cluster(name='rh-a10x', instance_type='A100:1')
- hf = SelfHostedHuggingFaceInstructEmbeddings(
- model_name=model_name, hardware=gpu)
- """ # noqa: E501
-
- model_id: str = DEFAULT_INSTRUCT_MODEL
- """Model name to use."""
- embed_instruction: str = DEFAULT_EMBED_INSTRUCTION
- """Instruction to use for embedding documents."""
- query_instruction: str = DEFAULT_QUERY_INSTRUCTION
- """Instruction to use for embedding query."""
- model_reqs: List[str] = ["./", "InstructorEmbedding", "torch"]
- """Requirements to install on hardware to inference the model."""
-
- def __init__(self, **kwargs: Any):
- """Initialize the remote inference function."""
- load_fn_kwargs = kwargs.pop("load_fn_kwargs", {})
- load_fn_kwargs["model_id"] = load_fn_kwargs.get(
- "model_id", DEFAULT_INSTRUCT_MODEL
- )
- load_fn_kwargs["instruct"] = load_fn_kwargs.get("instruct", True)
- load_fn_kwargs["device"] = load_fn_kwargs.get("device", 0)
- super().__init__(load_fn_kwargs=load_fn_kwargs, **kwargs)
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Compute doc embeddings using a HuggingFace instruct model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- instruction_pairs = []
- for text in texts:
- instruction_pairs.append([self.embed_instruction, text])
- embeddings = self.client(self.pipeline_ref, instruction_pairs)
- return embeddings.tolist()
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using a HuggingFace instruct model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- instruction_pair = [self.query_instruction, text]
- embedding = self.client(self.pipeline_ref, [instruction_pair])[0]
- return embedding.tolist()
diff --git a/libs/community/langchain_community/embeddings/sentence_transformer.py b/libs/community/langchain_community/embeddings/sentence_transformer.py
deleted file mode 100644
index 8a6d25e6c3..0000000000
--- a/libs/community/langchain_community/embeddings/sentence_transformer.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""HuggingFace sentence_transformer embedding models."""
-
-from langchain_community.embeddings.huggingface import HuggingFaceEmbeddings
-
-SentenceTransformerEmbeddings = HuggingFaceEmbeddings
diff --git a/libs/community/langchain_community/embeddings/solar.py b/libs/community/langchain_community/embeddings/solar.py
deleted file mode 100644
index 13d0c02648..0000000000
--- a/libs/community/langchain_community/embeddings/solar.py
+++ /dev/null
@@ -1,142 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Callable, Dict, List, Optional
-
-import requests
-from langchain_core._api import deprecated
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict, SecretStr
-from tenacity import (
- before_sleep_log,
- retry,
- stop_after_attempt,
- wait_exponential,
-)
-
-logger = logging.getLogger(__name__)
-
-
-def _create_retry_decorator() -> Callable[[Any], Any]:
- """Returns a tenacity retry decorator."""
-
- multiplier = 1
- min_seconds = 1
- max_seconds = 4
- max_retries = 6
-
- return retry(
- reraise=True,
- stop=stop_after_attempt(max_retries),
- wait=wait_exponential(multiplier=multiplier, min=min_seconds, max=max_seconds),
- before_sleep=before_sleep_log(logger, logging.WARNING),
- )
-
-
-def embed_with_retry(embeddings: SolarEmbeddings, *args: Any, **kwargs: Any) -> Any:
- """Use tenacity to retry the completion call."""
- retry_decorator = _create_retry_decorator()
-
- @retry_decorator
- def _embed_with_retry(*args: Any, **kwargs: Any) -> Any:
- return embeddings.embed(*args, **kwargs)
-
- return _embed_with_retry(*args, **kwargs)
-
-
-@deprecated(
- since="0.0.34", removal="1.0", alternative_import="langchain_upstage.ChatUpstage"
-)
-class SolarEmbeddings(BaseModel, Embeddings):
- """Solar's embedding service.
-
- To use, you should have the environment variable``SOLAR_API_KEY`` set
- with your API token, or pass it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import SolarEmbeddings
- embeddings = SolarEmbeddings()
-
- query_text = "This is a test query."
- query_result = embeddings.embed_query(query_text)
-
- document_text = "This is a test document."
- document_result = embeddings.embed_documents([document_text])
-
- """
-
- endpoint_url: str = "https://api.upstage.ai/v1/solar/embeddings"
- """Endpoint URL to use."""
- model: str = "embedding-query"
- """Embeddings model name to use."""
- solar_api_key: Optional[SecretStr] = None
- """API Key for Solar API."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate api key exists in environment."""
- solar_api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "solar_api_key", "SOLAR_API_KEY")
- )
- values["solar_api_key"] = solar_api_key
- return values
-
- def embed(
- self,
- text: str,
- ) -> List[List[float]]:
- payload = {
- "model": self.model,
- "input": text,
- }
-
- # HTTP headers for authorization
- headers = {
- "Authorization": f"Bearer {self.solar_api_key.get_secret_value()}", # type: ignore[union-attr]
- "Content-Type": "application/json",
- }
-
- # send request
- response = requests.post(self.endpoint_url, headers=headers, json=payload)
- parsed_response = response.json()
-
- # check for errors
- if len(parsed_response["data"]) == 0:
- raise ValueError(
- f"Solar API returned an error: {parsed_response['base_resp']}"
- )
-
- embedding = parsed_response["data"][0]["embedding"]
-
- return embedding
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed documents using a Solar embedding endpoint.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- embeddings = [embed_with_retry(self, text=text) for text in texts]
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using a Solar embedding endpoint.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- embedding = embed_with_retry(self, text=text)
- return embedding
diff --git a/libs/community/langchain_community/embeddings/spacy_embeddings.py b/libs/community/langchain_community/embeddings/spacy_embeddings.py
deleted file mode 100644
index cbe9c06d57..0000000000
--- a/libs/community/langchain_community/embeddings/spacy_embeddings.py
+++ /dev/null
@@ -1,116 +0,0 @@
-import importlib.util
-from typing import Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict, model_validator
-
-
-class SpacyEmbeddings(BaseModel, Embeddings):
- """Embeddings by spaCy models.
-
- Attributes:
- model_name (str): Name of a spaCy model.
- nlp (Any): The spaCy model loaded into memory.
-
- Methods:
- embed_documents(texts: List[str]) -> List[List[float]]:
- Generates embeddings for a list of documents.
- embed_query(text: str) -> List[float]:
- Generates an embedding for a single piece of text.
- """
-
- model_name: str = "en_core_web_sm"
- nlp: Optional[Any] = None
-
- model_config = ConfigDict(extra="forbid", protected_namespaces=())
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """
- Validates that the spaCy package and the model are installed.
-
- Args:
- values (Dict): The values provided to the class constructor.
-
- Returns:
- The validated values.
-
- Raises:
- ValueError: If the spaCy package or the
- model are not installed.
- """
- if values.get("model_name") is None:
- values["model_name"] = "en_core_web_sm"
-
- model_name = values.get("model_name")
-
- # Check if the spaCy package is installed
- if importlib.util.find_spec("spacy") is None:
- raise ValueError(
- "SpaCy package not found. Please install it with `pip install spacy`."
- )
- try:
- # Try to load the spaCy model
- import spacy
-
- values["nlp"] = spacy.load(model_name)
- except OSError:
- # If the model is not found, raise a ValueError
- raise ValueError(
- f"SpaCy model '{model_name}' not found. "
- f"Please install it with"
- f" `python -m spacy download {model_name}`"
- "or provide a valid spaCy model name."
- )
- return values # Return the validated values
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """
- Generates embeddings for a list of documents.
-
- Args:
- texts (List[str]): The documents to generate embeddings for.
-
- Returns:
- A list of embeddings, one for each document.
- """
- return [self.nlp(text).vector.tolist() for text in texts] # type: ignore[misc]
-
- def embed_query(self, text: str) -> List[float]:
- """
- Generates an embedding for a single piece of text.
-
- Args:
- text (str): The text to generate an embedding for.
-
- Returns:
- The embedding for the text.
- """
- return self.nlp(text).vector.tolist() # type: ignore[misc]
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- """
- Asynchronously generates embeddings for a list of documents.
- This method is not implemented and raises a NotImplementedError.
-
- Args:
- texts (List[str]): The documents to generate embeddings for.
-
- Raises:
- NotImplementedError: This method is not implemented.
- """
- raise NotImplementedError("Asynchronous embedding generation is not supported.")
-
- async def aembed_query(self, text: str) -> List[float]:
- """
- Asynchronously generates an embedding for a single piece of text.
- This method is not implemented and raises a NotImplementedError.
-
- Args:
- text (str): The text to generate an embedding for.
-
- Raises:
- NotImplementedError: This method is not implemented.
- """
- raise NotImplementedError("Asynchronous embedding generation is not supported.")
diff --git a/libs/community/langchain_community/embeddings/sparkllm.py b/libs/community/langchain_community/embeddings/sparkllm.py
deleted file mode 100644
index 6f0f1a1055..0000000000
--- a/libs/community/langchain_community/embeddings/sparkllm.py
+++ /dev/null
@@ -1,276 +0,0 @@
-import base64
-import hashlib
-import hmac
-import json
-import logging
-from datetime import datetime
-from time import mktime
-from typing import Any, Dict, List, Literal, Optional
-from urllib.parse import urlencode
-from wsgiref.handlers import format_date_time
-
-import numpy as np
-import requests
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import (
- secret_from_env,
-)
-from numpy import ndarray
-from pydantic import BaseModel, ConfigDict, Field, SecretStr
-
-# SparkLLMTextEmbeddings is an embedding model provided by iFLYTEK Co., Ltd.. (https://iflytek.com/en/).
-
-# Official Website: https://www.xfyun.cn/doc/spark/Embedding_api.html
-# Developers need to create an application in the console first, use the appid, APIKey,
-# and APISecret provided in the application for authentication,
-# and generate an authentication URL for handshake.
-# You can get one by registering at https://console.xfyun.cn/services/bm3.
-# SparkLLMTextEmbeddings support 2K token window and preduces vectors with
-# 2560 dimensions.
-
-logger = logging.getLogger(__name__)
-
-
-class Url:
- """URL class for parsing the URL."""
-
- def __init__(self, host: str, path: str, schema: str) -> None:
- self.host = host
- self.path = path
- self.schema = schema
- pass
-
-
-class SparkLLMTextEmbeddings(BaseModel, Embeddings):
- """SparkLLM embedding model integration.
-
- Setup:
- To use, you should have the environment variable "SPARK_APP_ID","SPARK_API_KEY"
- and "SPARK_API_SECRET" set your APP_ID, API_KEY and API_SECRET or pass it
- as a name parameter to the constructor.
-
- .. code-block:: bash
-
- export SPARK_APP_ID="your-api-id"
- export SPARK_API_KEY="your-api-key"
- export SPARK_API_SECRET="your-api-secret"
-
- Key init args — completion params:
- api_key: Optional[str]
- Automatically inferred from env var `SPARK_API_KEY` if not provided.
- app_id: Optional[str]
- Automatically inferred from env var `SPARK_APP_ID` if not provided.
- api_secret: Optional[str]
- Automatically inferred from env var `SPARK_API_SECRET` if not provided.
- base_url: Optional[str]
- Base URL path for API requests.
-
- See full list of supported init args and their descriptions in the params section.
-
- Instantiate:
-
- .. code-block:: python
-
- from langchain_community.embeddings import SparkLLMTextEmbeddings
-
- embed = SparkLLMTextEmbeddings(
- api_key="...",
- app_id="...",
- api_secret="...",
- # other
- )
-
- Embed single text:
- .. code-block:: python
-
- input_text = "The meaning of life is 42"
- embed.embed_query(input_text)
-
- .. code-block:: python
-
- [-0.4912109375, 0.60595703125, 0.658203125, 0.3037109375, 0.6591796875, 0.60302734375, ...]
-
- Embed multiple text:
- .. code-block:: python
-
- input_texts = ["This is a test query1.", "This is a test query2."]
- embed.embed_documents(input_texts)
-
- .. code-block:: python
-
- [
- [-0.1962890625, 0.94677734375, 0.7998046875, -0.1971435546875, 0.445556640625, 0.54638671875, ...],
- [ -0.44970703125, 0.06585693359375, 0.7421875, -0.474609375, 0.62353515625, 1.0478515625, ...],
- ]
- """ # noqa: E501
-
- spark_app_id: SecretStr = Field(
- alias="app_id", default_factory=secret_from_env("SPARK_APP_ID")
- )
- """Automatically inferred from env var `SPARK_APP_ID` if not provided."""
- spark_api_key: Optional[SecretStr] = Field(
- alias="api_key", default_factory=secret_from_env("SPARK_API_KEY", default=None)
- )
- """Automatically inferred from env var `SPARK_API_KEY` if not provided."""
- spark_api_secret: Optional[SecretStr] = Field(
- alias="api_secret",
- default_factory=secret_from_env("SPARK_API_SECRET", default=None),
- )
- """Automatically inferred from env var `SPARK_API_SECRET` if not provided."""
- base_url: str = Field(default="https://emb-cn-huabei-1.xf-yun.com/")
- """Base URL path for API requests"""
- domain: Literal["para", "query"] = Field(default="para")
- """This parameter is used for which Embedding this time belongs to.
- If "para"(default), it belongs to document Embedding.
- If "query", it belongs to query Embedding."""
-
- model_config = ConfigDict(
- populate_by_name=True,
- )
-
- def _embed(self, texts: List[str], host: str) -> Optional[List[List[float]]]:
- """Internal method to call Spark Embedding API and return embeddings.
-
- Args:
- texts: A list of texts to embed.
- host: Base URL path for API requests
-
- Returns:
- A list of list of floats representing the embeddings,
- or list with value None if an error occurs.
- """
- app_id = ""
- api_key = ""
- api_secret = ""
- if self.spark_app_id:
- app_id = self.spark_app_id.get_secret_value()
- if self.spark_api_key:
- api_key = self.spark_api_key.get_secret_value()
- if self.spark_api_secret:
- api_secret = self.spark_api_secret.get_secret_value()
- url = self._assemble_ws_auth_url(
- request_url=host,
- method="POST",
- api_key=api_key,
- api_secret=api_secret,
- )
- embed_result: list = []
- for text in texts:
- query_context = {"messages": [{"content": text, "role": "user"}]}
- content = self._get_body(app_id, query_context)
- response = requests.post(
- url, json=content, headers={"content-type": "application/json"}
- ).text
- res_arr = self._parser_message(response)
- if res_arr is not None:
- embed_result.append(res_arr.tolist())
- else:
- embed_result.append(None)
- return embed_result
-
- def embed_documents(self, texts: List[str]) -> Optional[List[List[float]]]: # type: ignore[override]
- """Public method to get embeddings for a list of documents.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- A list of embeddings, one for each text, or None if an error occurs.
- """
- return self._embed(texts, self.base_url)
-
- def embed_query(self, text: str) -> Optional[List[float]]: # type: ignore[override]
- """Public method to get embedding for a single query text.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text, or None if an error occurs.
- """
- result = self._embed([text], self.base_url)
- return result[0] if result is not None else None
-
- @staticmethod
- def _assemble_ws_auth_url(
- request_url: str, method: str = "GET", api_key: str = "", api_secret: str = ""
- ) -> str:
- u = SparkLLMTextEmbeddings._parse_url(request_url)
- host = u.host
- path = u.path
- now = datetime.now()
- date = format_date_time(mktime(now.timetuple()))
- signature_origin = "host: {}\ndate: {}\n{} {} HTTP/1.1".format(
- host, date, method, path
- )
- signature_sha = hmac.new(
- api_secret.encode("utf-8"),
- signature_origin.encode("utf-8"),
- digestmod=hashlib.sha256,
- ).digest()
- signature_sha_str = base64.b64encode(signature_sha).decode(encoding="utf-8")
- authorization_origin = (
- 'api_key="%s", algorithm="%s", headers="%s", signature="%s"'
- % (api_key, "hmac-sha256", "host date request-line", signature_sha_str)
- )
- authorization = base64.b64encode(authorization_origin.encode("utf-8")).decode(
- encoding="utf-8"
- )
- values = {"host": host, "date": date, "authorization": authorization}
-
- return request_url + "?" + urlencode(values)
-
- @staticmethod
- def _parse_url(request_url: str) -> Url:
- stidx = request_url.index("://")
- host = request_url[stidx + 3 :]
- schema = request_url[: stidx + 3]
- edidx = host.index("/")
- if edidx <= 0:
- raise AssembleHeaderException("invalid request url:" + request_url)
- path = host[edidx:]
- host = host[:edidx]
- u = Url(host, path, schema)
- return u
-
- def _get_body(self, appid: str, text: dict) -> Dict[str, Any]:
- body = {
- "header": {"app_id": appid, "uid": "39769795890", "status": 3},
- "parameter": {
- "emb": {"domain": self.domain, "feature": {"encoding": "utf8"}}
- },
- "payload": {
- "messages": {
- "text": base64.b64encode(json.dumps(text).encode("utf-8")).decode()
- }
- },
- }
- return body
-
- @staticmethod
- def _parser_message(
- message: str,
- ) -> Optional[ndarray]:
- data = json.loads(message)
- code = data["header"]["code"]
- if code != 0:
- logger.warning(f"Request error: {code}, {data}")
- return None
- else:
- text_base = data["payload"]["feature"]["text"]
- text_data = base64.b64decode(text_base)
- dt = np.dtype(np.float32)
- dt = dt.newbyteorder("<")
- text = np.frombuffer(text_data, dtype=dt)
- if len(text) > 2560:
- array = text[:2560]
- else:
- array = text
- return array
-
-
-class AssembleHeaderException(Exception):
- """Exception raised for errors in the header assembly."""
-
- def __init__(self, msg: str) -> None:
- self.message = msg
diff --git a/libs/community/langchain_community/embeddings/tensorflow_hub.py b/libs/community/langchain_community/embeddings/tensorflow_hub.py
deleted file mode 100644
index 270c9f17cb..0000000000
--- a/libs/community/langchain_community/embeddings/tensorflow_hub.py
+++ /dev/null
@@ -1,75 +0,0 @@
-from typing import Any, List
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict
-
-DEFAULT_MODEL_URL = "https://tfhub.dev/google/universal-sentence-encoder-multilingual/3"
-
-
-class TensorflowHubEmbeddings(BaseModel, Embeddings):
- """TensorflowHub embedding models.
-
- To use, you should have the ``tensorflow_text`` python package installed.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import TensorflowHubEmbeddings
- url = "https://tfhub.dev/google/universal-sentence-encoder-multilingual/3"
- tf = TensorflowHubEmbeddings(model_url=url)
- """
-
- embed: Any = None #: :meta private:
- model_url: str = DEFAULT_MODEL_URL
- """Model name to use."""
-
- def __init__(self, **kwargs: Any):
- """Initialize the tensorflow_hub and tensorflow_text."""
- super().__init__(**kwargs)
- try:
- import tensorflow_hub
- except ImportError:
- raise ImportError(
- "Could not import tensorflow-hub python package. "
- "Please install it with `pip install tensorflow-hub``."
- )
- try:
- import tensorflow_text # noqa
- except ImportError:
- raise ImportError(
- "Could not import tensorflow_text python package. "
- "Please install it with `pip install tensorflow_text``."
- )
-
- self.embed = tensorflow_hub.load(self.model_url)
-
- model_config = ConfigDict(
- extra="forbid",
- protected_namespaces=(),
- )
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Compute doc embeddings using a TensorflowHub embedding model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- texts = list(map(lambda x: x.replace("\n", " "), texts))
- embeddings = self.embed(texts).numpy()
- return embeddings.tolist()
-
- def embed_query(self, text: str) -> List[float]:
- """Compute query embeddings using a TensorflowHub embedding model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- text = text.replace("\n", " ")
- embedding = self.embed([text]).numpy()[0]
- return embedding.tolist()
diff --git a/libs/community/langchain_community/embeddings/text2vec.py b/libs/community/langchain_community/embeddings/text2vec.py
deleted file mode 100644
index 4b8cc77192..0000000000
--- a/libs/community/langchain_community/embeddings/text2vec.py
+++ /dev/null
@@ -1,81 +0,0 @@
-"""Wrapper around text2vec embedding models."""
-
-from typing import Any, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict
-
-
-class Text2vecEmbeddings(Embeddings, BaseModel):
- """text2vec embedding models.
-
- Install text2vec first, run 'pip install -U text2vec'.
- The github repository for text2vec is : https://github.com/shibing624/text2vec
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings.text2vec import Text2vecEmbeddings
-
- embedding = Text2vecEmbeddings()
- embedding.embed_documents([
- "This is a CoSENT(Cosine Sentence) model.",
- "It maps sentences to a 768 dimensional dense vector space.",
- ])
- embedding.embed_query(
- "It can be used for text matching or semantic search."
- )
- """
-
- model_name_or_path: Optional[str] = None
- encoder_type: Any = "MEAN"
- max_seq_length: int = 256
- device: Optional[str] = None
- model: Any = None
-
- model_config = ConfigDict(protected_namespaces=())
-
- def __init__(
- self,
- *,
- model: Any = None,
- model_name_or_path: Optional[str] = None,
- **kwargs: Any,
- ):
- try:
- from text2vec import SentenceModel
- except ImportError as e:
- raise ImportError(
- "Unable to import text2vec, please install with "
- "`pip install -U text2vec`."
- ) from e
-
- model_kwargs = {}
- if model_name_or_path is not None:
- model_kwargs["model_name_or_path"] = model_name_or_path
- model = model or SentenceModel(**model_kwargs, **kwargs)
- super().__init__(model=model, model_name_or_path=model_name_or_path, **kwargs)
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed documents using the text2vec embeddings model.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
-
- return self.model.encode(texts)
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using the text2vec embeddings model.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
-
- return self.model.encode(text)
diff --git a/libs/community/langchain_community/embeddings/textembed.py b/libs/community/langchain_community/embeddings/textembed.py
deleted file mode 100644
index 6d963b3d1b..0000000000
--- a/libs/community/langchain_community/embeddings/textembed.py
+++ /dev/null
@@ -1,350 +0,0 @@
-"""
-TextEmbed: Embedding Inference Server
-
-TextEmbed provides a high-throughput, low-latency solution for serving embeddings.
-It supports various sentence-transformer models.
-Now, it includes the ability to deploy image embedding models.
-TextEmbed offers flexibility and scalability for diverse applications.
-
-TextEmbed is maintained by Keval Dekivadiya and is licensed under the Apache-2.0 license.
-""" # noqa: E501
-
-import asyncio
-from concurrent.futures import ThreadPoolExecutor
-from typing import Any, Callable, Dict, List, Optional, Tuple, Union
-
-import aiohttp
-import numpy as np
-import requests
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import from_env, secret_from_env
-from pydantic import BaseModel, ConfigDict, Field, SecretStr, model_validator
-from typing_extensions import Self
-
-__all__ = ["TextEmbedEmbeddings"]
-
-
-class TextEmbedEmbeddings(BaseModel, Embeddings):
- """
- A class to handle embedding requests to the TextEmbed API.
-
- Attributes:
- model : The TextEmbed model ID to use for embeddings.
- api_url : The base URL for the TextEmbed API.
- api_key : The API key for authenticating with the TextEmbed API.
- client : The TextEmbed client instance.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import TextEmbedEmbeddings
-
- embeddings = TextEmbedEmbeddings(
- model="sentence-transformers/clip-ViT-B-32",
- api_url="http://localhost:8000/v1",
- api_key=""
- )
-
- For more information: https://github.com/kevaldekivadiya2415/textembed/blob/main/docs/setup.md
- """ # noqa: E501
-
- model: str
- """Underlying TextEmbed model id."""
-
- api_url: str = Field(
- default_factory=from_env(
- "TEXTEMBED_API_URL", default="http://localhost:8000/v1"
- )
- )
- """Endpoint URL to use."""
-
- api_key: SecretStr = Field(default_factory=secret_from_env("TEXTEMBED_API_KEY"))
- """API Key for authentication"""
-
- client: Any = None
- """TextEmbed client."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="after")
- def validate_environment(self) -> Self:
- """Validate that api key and URL exist in the environment."""
- self.client = AsyncOpenAITextEmbedEmbeddingClient(
- host=self.api_url, api_key=self.api_key.get_secret_value()
- )
- return self
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Call out to TextEmbed's embedding endpoint.
-
- Args:
- texts (List[str]): The list of texts to embed.
-
- Returns:
- List[List[float]]: List of embeddings, one for each text.
- """
- embeddings = self.client.embed(
- model=self.model,
- texts=texts,
- )
- return embeddings
-
- async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
- """Async call out to TextEmbed's embedding endpoint.
-
- Args:
- texts (List[str]): The list of texts to embed.
-
- Returns:
- List[List[float]]: List of embeddings, one for each text.
- """
- embeddings = await self.client.aembed(
- model=self.model,
- texts=texts,
- )
- return embeddings
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to TextEmbed's embedding endpoint for a single query.
-
- Args:
- text (str): The text to embed.
-
- Returns:
- List[float]: Embeddings for the text.
- """
- return self.embed_documents([text])[0]
-
- async def aembed_query(self, text: str) -> List[float]:
- """Async call out to TextEmbed's embedding endpoint for a single query.
-
- Args:
- text (str): The text to embed.
-
- Returns:
- List[float]: Embeddings for the text.
- """
- embeddings = await self.aembed_documents([text])
- return embeddings[0]
-
-
-class AsyncOpenAITextEmbedEmbeddingClient:
- """
- A client to handle synchronous and asynchronous requests to the TextEmbed API.
-
- Attributes:
- host (str): The base URL for the TextEmbed API.
- api_key (str): The API key for authenticating with the TextEmbed API.
- aiosession (Optional[aiohttp.ClientSession]): The aiohttp session for async requests.
- _batch_size (int): Maximum batch size for a single request.
- """ # noqa: E501
-
- def __init__(
- self,
- host: str = "http://localhost:8000/v1",
- api_key: Union[str, None] = None,
- aiosession: Optional[aiohttp.ClientSession] = None,
- ) -> None:
- self.host = host
- self.api_key = api_key
- self.aiosession = aiosession
-
- if self.host is None or len(self.host) < 3:
- raise ValueError("Parameter `host` must be set to a valid URL")
- self._batch_size = 256
-
- @staticmethod
- def _permute(
- texts: List[str], sorter: Callable = len
- ) -> Tuple[List[str], Callable]:
- """
- Sorts texts in ascending order and provides a function to restore the original order.
-
- Args:
- texts (List[str]): List of texts to sort.
- sorter (Callable, optional): Sorting function, defaults to length.
-
- Returns:
- Tuple[List[str], Callable]: Sorted texts and a function to restore original order.
- """ # noqa: E501
- if len(texts) == 1:
- return texts, lambda t: t
- length_sorted_idx = np.argsort([-sorter(sen) for sen in texts])
- texts_sorted = [texts[idx] for idx in length_sorted_idx]
-
- return texts_sorted, lambda unsorted_embeddings: [
- unsorted_embeddings[idx] for idx in np.argsort(length_sorted_idx)
- ]
-
- def _batch(self, texts: List[str]) -> List[List[str]]:
- """
- Splits a list of texts into batches of size max `self._batch_size`.
-
- Args:
- texts (List[str]): List of texts to split.
-
- Returns:
- List[List[str]]: List of batches of texts.
- """
- if len(texts) == 1:
- return [texts]
- batches = []
- for start_index in range(0, len(texts), self._batch_size):
- batches.append(texts[start_index : start_index + self._batch_size])
- return batches
-
- @staticmethod
- def _unbatch(batch_of_texts: List[List[Any]]) -> List[Any]:
- """
- Merges batches of texts into a single list.
-
- Args:
- batch_of_texts (List[List[Any]]): List of batches of texts.
-
- Returns:
- List[Any]: Merged list of texts.
- """
- if len(batch_of_texts) == 1 and len(batch_of_texts[0]) == 1:
- return batch_of_texts[0]
- texts = []
- for sublist in batch_of_texts:
- texts.extend(sublist)
- return texts
-
- def _kwargs_post_request(self, model: str, texts: List[str]) -> Dict[str, Any]:
- """
- Builds the kwargs for the POST request, used by sync method.
-
- Args:
- model (str): The model to use for embedding.
- texts (List[str]): List of texts to embed.
-
- Returns:
- Dict[str, Any]: Dictionary of POST request parameters.
- """
- return dict(
- url=f"{self.host}/embedding",
- headers={
- "accept": "application/json",
- "content-type": "application/json",
- "Authorization": f"Bearer {self.api_key}",
- },
- json=dict(
- input=texts,
- model=model,
- ),
- )
-
- def _sync_request_embed(
- self, model: str, batch_texts: List[str]
- ) -> List[List[float]]:
- """
- Sends a synchronous request to the embedding endpoint.
-
- Args:
- model (str): The model to use for embedding.
- batch_texts (List[str]): Batch of texts to embed.
-
- Returns:
- List[List[float]]: List of embeddings for the batch.
-
- Raises:
- Exception: If the response status is not 200.
- """
- response = requests.post(
- **self._kwargs_post_request(model=model, texts=batch_texts)
- )
- if response.status_code != 200:
- raise Exception(
- f"TextEmbed responded with an unexpected status message "
- f"{response.status_code}: {response.text}"
- )
- return [e["embedding"] for e in response.json()["data"]]
-
- def embed(self, model: str, texts: List[str]) -> List[List[float]]:
- """
- Embeds a list of texts synchronously.
-
- Args:
- model (str): The model to use for embedding.
- texts (List[str]): List of texts to embed.
-
- Returns:
- List[List[float]]: List of embeddings for the texts.
- """
- perm_texts, unpermute_func = self._permute(texts)
- perm_texts_batched = self._batch(perm_texts)
-
- # Request
- map_args = (
- self._sync_request_embed,
- [model] * len(perm_texts_batched),
- perm_texts_batched,
- )
- if len(perm_texts_batched) == 1:
- embeddings_batch_perm = list(map(*map_args))
- else:
- with ThreadPoolExecutor(32) as p:
- embeddings_batch_perm = list(p.map(*map_args))
-
- embeddings_perm = self._unbatch(embeddings_batch_perm)
- embeddings = unpermute_func(embeddings_perm)
- return embeddings
-
- async def _async_request(
- self, session: aiohttp.ClientSession, **kwargs: Dict[str, Any]
- ) -> List[List[float]]:
- """
- Sends an asynchronous request to the embedding endpoint.
-
- Args:
- session (aiohttp.ClientSession): The aiohttp session for the request.
- kwargs (Dict[str, Any]): Dictionary of POST request parameters.
-
- Returns:
- List[List[float]]: List of embeddings for the request.
-
- Raises:
- Exception: If the response status is not 200.
- """
- async with session.post(**kwargs) as response: # type: ignore[arg-type]
- if response.status != 200:
- raise Exception(
- f"TextEmbed responded with an unexpected status message "
- f"{response.status}: {response.text}"
- )
- embedding = (await response.json())["data"]
- return [e["embedding"] for e in embedding]
-
- async def aembed(self, model: str, texts: List[str]) -> List[List[float]]:
- """
- Embeds a list of texts asynchronously.
-
- Args:
- model (str): The model to use for embedding.
- texts (List[str]): List of texts to embed.
-
- Returns:
- List[List[float]]: List of embeddings for the texts.
- """
- perm_texts, unpermute_func = self._permute(texts)
- perm_texts_batched = self._batch(perm_texts)
-
- async with aiohttp.ClientSession(
- connector=aiohttp.TCPConnector(limit=32)
- ) as session:
- embeddings_batch_perm = await asyncio.gather(
- *[
- self._async_request(
- session=session,
- **self._kwargs_post_request(model=model, texts=t),
- )
- for t in perm_texts_batched
- ]
- )
-
- embeddings_perm = self._unbatch(embeddings_batch_perm)
- embeddings = unpermute_func(embeddings_perm)
- return embeddings
diff --git a/libs/community/langchain_community/embeddings/titan_takeoff.py b/libs/community/langchain_community/embeddings/titan_takeoff.py
deleted file mode 100644
index b171be0f29..0000000000
--- a/libs/community/langchain_community/embeddings/titan_takeoff.py
+++ /dev/null
@@ -1,210 +0,0 @@
-from enum import Enum
-from typing import Any, Dict, List, Optional, Set, Union
-
-from langchain_core.embeddings import Embeddings
-from pydantic import BaseModel, ConfigDict
-
-
-class TakeoffEmbeddingException(Exception):
- """Custom exception for interfacing with Takeoff Embedding class."""
-
-
-class MissingConsumerGroup(TakeoffEmbeddingException):
- """Exception raised when no consumer group is provided on initialization of
- TitanTakeoffEmbed or in embed request."""
-
-
-class Device(str, Enum):
- """Device to use for inference, cuda or cpu."""
-
- cuda = "cuda"
- cpu = "cpu"
-
-
-class ReaderConfig(BaseModel):
- """Configuration for the reader to be deployed in Takeoff."""
-
- model_config = ConfigDict(
- protected_namespaces=(),
- )
-
- model_name: str
- """The name of the model to use"""
-
- device: Device = Device.cuda
- """The device to use for inference, cuda or cpu"""
-
- consumer_group: str = "primary"
- """The consumer group to place the reader into"""
-
-
-class TitanTakeoffEmbed(Embeddings):
- """Interface with Takeoff Inference API for embedding models.
-
- Use it to send embedding requests and to deploy embedding
- readers with Takeoff.
-
- Examples:
- This is an example how to deploy an embedding model and send requests.
-
- .. code-block:: python
- # Import the TitanTakeoffEmbed class from community package
- import time
- from langchain_community.embeddings import TitanTakeoffEmbed
-
- # Specify the embedding reader you'd like to deploy
- reader_1 = {
- "model_name": "avsolatorio/GIST-large-Embedding-v0",
- "device": "cpu",
- "consumer_group": "embed"
- }
-
- # For every reader you pass into models arg Takeoff will spin up a reader
- # according to the specs you provide. If you don't specify the arg no models
- # are spun up and it assumes you have already done this separately.
- embed = TitanTakeoffEmbed(models=[reader_1])
-
- # Wait for the reader to be deployed, time needed depends on the model size
- # and your internet speed
- time.sleep(60)
-
- # Returns the embedded query, ie a List[float], sent to `embed` consumer
- # group where we just spun up the embedding reader
- print(embed.embed_query(
- "Where can I see football?", consumer_group="embed"
- ))
-
- # Returns a List of embeddings, ie a List[List[float]], sent to `embed`
- # consumer group where we just spun up the embedding reader
- print(embed.embed_document(
- ["Document1", "Document2"],
- consumer_group="embed"
- ))
- """
-
- base_url: str = "http://localhost"
- """The base URL of the Titan Takeoff (Pro) server. Default = "http://localhost"."""
-
- port: int = 3000
- """The port of the Titan Takeoff (Pro) server. Default = 3000."""
-
- mgmt_port: int = 3001
- """The management port of the Titan Takeoff (Pro) server. Default = 3001."""
-
- client: Any = None
- """Takeoff Client Python SDK used to interact with Takeoff API"""
-
- embed_consumer_groups: Set[str] = set()
- """The consumer groups in Takeoff which contain embedding models"""
-
- def __init__(
- self,
- base_url: str = "http://localhost",
- port: int = 3000,
- mgmt_port: int = 3001,
- models: List[ReaderConfig] = [],
- ):
- """Initialize the Titan Takeoff embedding wrapper.
-
- Args:
- base_url (str, optional): The base url where Takeoff Inference Server is
- listening. Defaults to "http://localhost".
- port (int, optional): What port is Takeoff Inference API listening on.
- Defaults to 3000.
- mgmt_port (int, optional): What port is Takeoff Management API listening on.
- Defaults to 3001.
- models (List[ReaderConfig], optional): Any readers you'd like to spin up on.
- Defaults to [].
-
- Raises:
- ImportError: If you haven't installed takeoff-client, you will get an
- ImportError. To remedy run `pip install 'takeoff-client==0.4.0'`
- """
- self.base_url = base_url
- self.port = port
- self.mgmt_port = mgmt_port
- try:
- from takeoff_client import TakeoffClient
- except ImportError:
- raise ImportError(
- "takeoff-client is required for TitanTakeoff. "
- "Please install it with `pip install 'takeoff-client==0.4.0'`."
- )
- self.client = TakeoffClient(
- self.base_url, port=self.port, mgmt_port=self.mgmt_port
- )
- for model in models:
- self.client.create_reader(model)
- if isinstance(model, dict):
- self.embed_consumer_groups.add(model.get("consumer_group"))
- else:
- self.embed_consumer_groups.add(model.consumer_group)
- super(TitanTakeoffEmbed, self).__init__()
-
- def _embed(
- self, input: Union[List[str], str], consumer_group: Optional[str]
- ) -> Dict[str, Any]:
- """Embed text.
-
- Args:
- input (Union[List[str], str]): prompt/document or list of prompts/documents
- to embed
- consumer_group (Optional[str]): what consumer group to send the embedding
- request to. If not specified and there is only one
- consumer group specified during initialization, it will be used. If there
- are multiple consumer groups specified during initialization, you must
- specify which one to use.
-
- Raises:
- MissingConsumerGroup: The consumer group can not be inferred from the
- initialization and must be specified with request.
-
- Returns:
- Dict[str, Any]: Result of query, {"result": List[List[float]]} or
- {"result": List[float]}
- """
- if not consumer_group:
- if len(self.embed_consumer_groups) == 1:
- consumer_group = list(self.embed_consumer_groups)[0]
- elif len(self.embed_consumer_groups) > 1:
- raise MissingConsumerGroup(
- "TakeoffEmbedding was initialized with multiple embedding reader"
- "groups, you must specify which one to use."
- )
- else:
- raise MissingConsumerGroup(
- "You must specify what consumer group you want to send embedding"
- "response to as TitanTakeoffEmbed was not initialized with an "
- "embedding reader."
- )
- return self.client.embed(input, consumer_group)
-
- def embed_documents(
- self, texts: List[str], consumer_group: Optional[str] = None
- ) -> List[List[float]]:
- """Embed documents.
-
- Args:
- texts (List[str]): List of prompts/documents to embed
- consumer_group (Optional[str], optional): Consumer group to send request
- to containing embedding model. Defaults to None.
-
- Returns:
- List[List[float]]: List of embeddings
- """
- return self._embed(texts, consumer_group)["result"]
-
- def embed_query(
- self, text: str, consumer_group: Optional[str] = None
- ) -> List[float]:
- """Embed query.
-
- Args:
- text (str): Prompt/document to embed
- consumer_group (Optional[str], optional): Consumer group to send request
- to containing embedding model. Defaults to None.
-
- Returns:
- List[float]: Embedding
- """
- return self._embed(text, consumer_group)["result"]
diff --git a/libs/community/langchain_community/embeddings/vertexai.py b/libs/community/langchain_community/embeddings/vertexai.py
deleted file mode 100644
index 06637385e0..0000000000
--- a/libs/community/langchain_community/embeddings/vertexai.py
+++ /dev/null
@@ -1,361 +0,0 @@
-import logging
-import re
-import string
-import threading
-from concurrent.futures import ThreadPoolExecutor, wait
-from typing import Any, Dict, List, Literal, Optional, Tuple
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.embeddings import Embeddings
-from langchain_core.language_models.llms import create_base_retry_decorator
-from langchain_core.utils import pre_init
-
-from langchain_community.llms.vertexai import _VertexAICommon
-from langchain_community.utilities.vertexai import raise_vertex_import_error
-
-logger = logging.getLogger(__name__)
-
-_MAX_TOKENS_PER_BATCH = 20000
-_MAX_BATCH_SIZE = 250
-_MIN_BATCH_SIZE = 5
-
-
-@deprecated(
- since="0.0.12",
- removal="1.0",
- alternative_import="langchain_google_vertexai.VertexAIEmbeddings",
-)
-class VertexAIEmbeddings(_VertexAICommon, Embeddings):
- """Google Cloud VertexAI embedding models."""
-
- # Instance context
- instance: Dict[str, Any] = {} #: :meta private:
- show_progress_bar: bool = False
- """Whether to show a tqdm progress bar. Must have `tqdm` installed."""
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validates that the python package exists in environment."""
- cls._try_init_vertexai(values)
- if values["model_name"] == "textembedding-gecko-default":
- logger.warning(
- "Model_name will become a required arg for VertexAIEmbeddings "
- "starting from Feb-01-2024. Currently the default is set to "
- "textembedding-gecko@001"
- )
- values["model_name"] = "textembedding-gecko@001"
- try:
- from vertexai.language_models import TextEmbeddingModel
- except ImportError:
- raise_vertex_import_error()
- values["client"] = TextEmbeddingModel.from_pretrained(values["model_name"])
- return values
-
- def __init__(
- self,
- # the default value would be removed after Feb-01-2024
- model_name: str = "textembedding-gecko-default",
- project: Optional[str] = None,
- location: str = "us-central1",
- request_parallelism: int = 5,
- max_retries: int = 6,
- credentials: Optional[Any] = None,
- **kwargs: Any,
- ):
- """Initialize the sentence_transformer."""
- super().__init__(
- project=project,
- location=location,
- credentials=credentials,
- request_parallelism=request_parallelism,
- max_retries=max_retries,
- model_name=model_name,
- **kwargs,
- )
- self.instance["max_batch_size"] = kwargs.get("max_batch_size", _MAX_BATCH_SIZE)
- self.instance["batch_size"] = self.instance["max_batch_size"]
- self.instance["min_batch_size"] = kwargs.get("min_batch_size", _MIN_BATCH_SIZE)
- self.instance["min_good_batch_size"] = self.instance["min_batch_size"]
- self.instance["lock"] = threading.Lock()
- self.instance["batch_size_validated"] = False
- self.instance["task_executor"] = ThreadPoolExecutor(
- max_workers=request_parallelism
- )
- self.instance[
- "embeddings_task_type_supported"
- ] = not self.client._endpoint_name.endswith("/textembedding-gecko@001")
-
- @staticmethod
- def _split_by_punctuation(text: str) -> List[str]:
- """Splits a string by punctuation and whitespace characters."""
- split_by = string.punctuation + "\t\n "
- pattern = f"([{split_by}])"
- # Using re.split to split the text based on the pattern
- return [segment for segment in re.split(pattern, text) if segment]
-
- @staticmethod
- def _prepare_batches(texts: List[str], batch_size: int) -> List[List[str]]:
- """Splits texts in batches based on current maximum batch size
- and maximum tokens per request.
- """
- text_index = 0
- texts_len = len(texts)
- batch_token_len = 0
- batches: List[List[str]] = []
- current_batch: List[str] = []
- if texts_len == 0:
- return []
- while text_index < texts_len:
- current_text = texts[text_index]
- # Number of tokens per a text is conservatively estimated
- # as 2 times number of words, punctuation and whitespace characters.
- # Using `count_tokens` API will make batching too expensive.
- # Utilizing a tokenizer, would add a dependency that would not
- # necessarily be reused by the application using this class.
- current_text_token_cnt = (
- len(VertexAIEmbeddings._split_by_punctuation(current_text)) * 2
- )
- end_of_batch = False
- if current_text_token_cnt > _MAX_TOKENS_PER_BATCH:
- # Current text is too big even for a single batch.
- # Such request will fail, but we still make a batch
- # so that the app can get the error from the API.
- if len(current_batch) > 0:
- # Adding current batch if not empty.
- batches.append(current_batch)
- current_batch = [current_text]
- text_index += 1
- end_of_batch = True
- elif (
- batch_token_len + current_text_token_cnt > _MAX_TOKENS_PER_BATCH
- or len(current_batch) == batch_size
- ):
- end_of_batch = True
- else:
- if text_index == texts_len - 1:
- # Last element - even though the batch may be not big,
- # we still need to make it.
- end_of_batch = True
- batch_token_len += current_text_token_cnt
- current_batch.append(current_text)
- text_index += 1
- if end_of_batch:
- batches.append(current_batch)
- current_batch = []
- batch_token_len = 0
- return batches
-
- def _get_embeddings_with_retry(
- self, texts: List[str], embeddings_type: Optional[str] = None
- ) -> List[List[float]]:
- """Makes a Vertex AI model request with retry logic."""
- from google.api_core.exceptions import (
- Aborted,
- DeadlineExceeded,
- ResourceExhausted,
- ServiceUnavailable,
- )
-
- errors = [
- ResourceExhausted,
- ServiceUnavailable,
- Aborted,
- DeadlineExceeded,
- ]
- retry_decorator = create_base_retry_decorator(
- error_types=errors,
- max_retries=self.max_retries,
- )
-
- @retry_decorator
- def _completion_with_retry(texts_to_process: List[str]) -> Any:
- if embeddings_type and self.instance["embeddings_task_type_supported"]:
- from vertexai.language_models import TextEmbeddingInput
-
- requests = [
- TextEmbeddingInput(text=t, task_type=embeddings_type)
- for t in texts_to_process
- ]
- else:
- requests = texts_to_process
- embeddings = self.client.get_embeddings(requests)
- return [embs.values for embs in embeddings]
-
- return _completion_with_retry(texts)
-
- def _prepare_and_validate_batches(
- self, texts: List[str], embeddings_type: Optional[str] = None
- ) -> Tuple[List[List[float]], List[List[str]]]:
- """Prepares text batches with one-time validation of batch size.
- Batch size varies between GCP regions and individual project quotas.
- # Returns embeddings of the first text batch that went through,
- # and text batches for the rest of the texts.
- """
- from google.api_core.exceptions import InvalidArgument
-
- batches = VertexAIEmbeddings._prepare_batches(
- texts, self.instance["batch_size"]
- )
- # If batch size if less or equal to one that went through before,
- # then keep batches as they are.
- if len(batches[0]) <= self.instance["min_good_batch_size"]:
- return [], batches
- with self.instance["lock"]:
- # If largest possible batch size was validated
- # while waiting for the lock, then check for rebuilding
- # our batches, and return.
- if self.instance["batch_size_validated"]:
- if len(batches[0]) <= self.instance["batch_size"]:
- return [], batches
- else:
- return [], VertexAIEmbeddings._prepare_batches(
- texts, self.instance["batch_size"]
- )
- # Figure out largest possible batch size by trying to push
- # batches and lowering their size in half after every failure.
- first_batch = batches[0]
- first_result = []
- had_failure = False
- while True:
- try:
- first_result = self._get_embeddings_with_retry(
- first_batch, embeddings_type
- )
- break
- except InvalidArgument:
- had_failure = True
- first_batch_len = len(first_batch)
- if first_batch_len == self.instance["min_batch_size"]:
- raise
- first_batch_len = max(
- self.instance["min_batch_size"], int(first_batch_len / 2)
- )
- first_batch = first_batch[:first_batch_len]
- first_batch_len = len(first_batch)
- self.instance["min_good_batch_size"] = max(
- self.instance["min_good_batch_size"], first_batch_len
- )
- # If had a failure and recovered
- # or went through with the max size, then it's a legit batch size.
- if had_failure or first_batch_len == self.instance["max_batch_size"]:
- self.instance["batch_size"] = first_batch_len
- self.instance["batch_size_validated"] = True
- # If batch size was updated,
- # rebuild batches with the new batch size
- # (texts that went through are excluded here).
- if first_batch_len != self.instance["max_batch_size"]:
- batches = VertexAIEmbeddings._prepare_batches(
- texts[first_batch_len:], self.instance["batch_size"]
- )
- else:
- # Still figuring out max batch size.
- batches = batches[1:]
- # Returning embeddings of the first text batch that went through,
- # and text batches for the rest of texts.
- return first_result, batches
-
- def embed(
- self,
- texts: List[str],
- batch_size: int = 0,
- embeddings_task_type: Optional[
- Literal[
- "RETRIEVAL_QUERY",
- "RETRIEVAL_DOCUMENT",
- "SEMANTIC_SIMILARITY",
- "CLASSIFICATION",
- "CLUSTERING",
- ]
- ] = None,
- ) -> List[List[float]]:
- """Embed a list of strings.
-
- Args:
- texts: List[str] The list of strings to embed.
- batch_size: [int] The batch size of embeddings to send to the model.
- If zero, then the largest batch size will be detected dynamically
- at the first request, starting from 250, down to 5.
- embeddings_task_type: [str] optional embeddings task type,
- one of the following
- RETRIEVAL_QUERY - Text is a query
- in a search/retrieval setting.
- RETRIEVAL_DOCUMENT - Text is a document
- in a search/retrieval setting.
- SEMANTIC_SIMILARITY - Embeddings will be used
- for Semantic Textual Similarity (STS).
- CLASSIFICATION - Embeddings will be used for classification.
- CLUSTERING - Embeddings will be used for clustering.
-
- Returns:
- List of embeddings, one for each text.
- """
- if len(texts) == 0:
- return []
- embeddings: List[List[float]] = []
- first_batch_result: List[List[float]] = []
- if batch_size > 0:
- # Fixed batch size.
- batches = VertexAIEmbeddings._prepare_batches(texts, batch_size)
- else:
- # Dynamic batch size, starting from 250 at the first call.
- first_batch_result, batches = self._prepare_and_validate_batches(
- texts, embeddings_task_type
- )
- # First batch result may have some embeddings already.
- # In such case, batches have texts that were not processed yet.
- embeddings.extend(first_batch_result)
- tasks = []
- if self.show_progress_bar:
- try:
- from tqdm import tqdm
-
- iter_ = tqdm(batches, desc="VertexAIEmbeddings")
- except ImportError:
- logger.warning(
- "Unable to show progress bar because tqdm could not be imported. "
- "Please install with `pip install tqdm`."
- )
- iter_ = batches
- else:
- iter_ = batches
- for batch in iter_:
- tasks.append(
- self.instance["task_executor"].submit(
- self._get_embeddings_with_retry,
- texts=batch,
- embeddings_type=embeddings_task_type,
- )
- )
- if len(tasks) > 0:
- wait(tasks)
- for t in tasks:
- embeddings.extend(t.result())
- return embeddings
-
- def embed_documents(
- self, texts: List[str], batch_size: int = 0
- ) -> List[List[float]]:
- """Embed a list of documents.
-
- Args:
- texts: List[str] The list of texts to embed.
- batch_size: [int] The batch size of embeddings to send to the model.
- If zero, then the largest batch size will be detected dynamically
- at the first request, starting from 250, down to 5.
-
- Returns:
- List of embeddings, one for each text.
- """
- return self.embed(texts, batch_size, "RETRIEVAL_DOCUMENT")
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a text.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embedding for the text.
- """
- embeddings = self.embed([text], 1, "RETRIEVAL_QUERY")
- return embeddings[0]
diff --git a/libs/community/langchain_community/embeddings/volcengine.py b/libs/community/langchain_community/embeddings/volcengine.py
deleted file mode 100644
index 6417e5e5a3..0000000000
--- a/libs/community/langchain_community/embeddings/volcengine.py
+++ /dev/null
@@ -1,128 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import BaseModel
-
-logger = logging.getLogger(__name__)
-
-
-class VolcanoEmbeddings(BaseModel, Embeddings):
- """`Volcengine Embeddings` embedding models."""
-
- volcano_ak: Optional[str] = None
- """volcano access key
- learn more from: https://www.volcengine.com/docs/6459/76491#ak-sk"""
-
- volcano_sk: Optional[str] = None
- """volcano secret key
- learn more from: https://www.volcengine.com/docs/6459/76491#ak-sk"""
-
- host: str = "maas-api.ml-platform-cn-beijing.volces.com"
- """host
- learn more from https://www.volcengine.com/docs/82379/1174746"""
- region: str = "cn-beijing"
- """region
- learn more from https://www.volcengine.com/docs/82379/1174746"""
-
- model: str = "bge-large-zh"
- """Model name
- you could get from https://www.volcengine.com/docs/82379/1174746
- for now, we support bge_large_zh
- """
-
- version: str = "1.0"
- """ model version """
-
- chunk_size: int = 100
- """Chunk size when multiple texts are input"""
-
- client: Any
- """volcano client"""
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """
- Validate whether volcano_ak and volcano_sk in the environment variables or
- configuration file are available or not.
-
- init volcano embedding client with `ak`, `sk`, `host`, `region`
-
- Args:
-
- values: a dictionary containing configuration information, must include the
- fields of volcano_ak and volcano_sk
- Returns:
-
- a dictionary containing configuration information. If volcano_ak and
- volcano_sk are not provided in the environment variables or configuration
- file,the original values will be returned; otherwise, values containing
- volcano_ak and volcano_sk will be returned.
- Raises:
-
- ValueError: volcengine package not found, please install it with
- `pip install volcengine`
- """
- values["volcano_ak"] = get_from_dict_or_env(
- values,
- "volcano_ak",
- "VOLC_ACCESSKEY",
- )
- values["volcano_sk"] = get_from_dict_or_env(
- values,
- "volcano_sk",
- "VOLC_SECRETKEY",
- )
-
- try:
- from volcengine.maas import MaasService
-
- client = MaasService(values["host"], values["region"])
- client.set_ak(values["volcano_ak"])
- client.set_sk(values["volcano_sk"])
- values["client"] = client
- except ImportError:
- raise ImportError(
- "volcengine package not found, please install it with "
- "`pip install volcengine`"
- )
- return values
-
- def embed_query(self, text: str) -> List[float]:
- return self.embed_documents([text])[0]
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """
- Embeds a list of text documents using the AutoVOT algorithm.
-
- Args:
- texts (List[str]): A list of text documents to embed.
-
- Returns:
- List[List[float]]: A list of embeddings for each document in the input list.
- Each embedding is represented as a list of float values.
- """
- text_in_chunks = [
- texts[i : i + self.chunk_size]
- for i in range(0, len(texts), self.chunk_size)
- ]
- lst = []
- for chunk in text_in_chunks:
- req = {
- "model": {
- "name": self.model,
- "version": self.version,
- },
- "input": chunk,
- }
- try:
- from volcengine.maas import MaasException
-
- resp = self.client.embeddings(req)
- lst.extend([res["embedding"] for res in resp["data"]])
- except MaasException as e:
- raise ValueError(f"embed by volcengine Error: {e}")
- return lst
diff --git a/libs/community/langchain_community/embeddings/voyageai.py b/libs/community/langchain_community/embeddings/voyageai.py
deleted file mode 100644
index 2ef1477ac6..0000000000
--- a/libs/community/langchain_community/embeddings/voyageai.py
+++ /dev/null
@@ -1,230 +0,0 @@
-from __future__ import annotations
-
-import json
-import logging
-from typing import (
- Any,
- Callable,
- Dict,
- List,
- Optional,
- Tuple,
- Union,
- cast,
-)
-
-import requests
-from langchain_core._api.deprecation import deprecated
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env
-from pydantic import BaseModel, ConfigDict, SecretStr, model_validator
-from tenacity import (
- before_sleep_log,
- retry,
- stop_after_attempt,
- wait_exponential,
-)
-
-logger = logging.getLogger(__name__)
-
-
-def _create_retry_decorator(embeddings: VoyageEmbeddings) -> Callable[[Any], Any]:
- min_seconds = 4
- max_seconds = 10
- # Wait 2^x * 1 second between each retry starting with
- # 4 seconds, then up to 10 seconds, then 10 seconds afterwards
- return retry(
- reraise=True,
- stop=stop_after_attempt(embeddings.max_retries),
- wait=wait_exponential(multiplier=1, min=min_seconds, max=max_seconds),
- before_sleep=before_sleep_log(logger, logging.WARNING),
- )
-
-
-def _check_response(response: dict) -> dict:
- if "data" not in response:
- raise RuntimeError(f"Voyage API Error. Message: {json.dumps(response)}")
- return response
-
-
-def embed_with_retry(embeddings: VoyageEmbeddings, **kwargs: Any) -> Any:
- """Use tenacity to retry the embedding call."""
- retry_decorator = _create_retry_decorator(embeddings)
-
- @retry_decorator
- def _embed_with_retry(**kwargs: Any) -> Any:
- response = requests.post(**kwargs)
- return _check_response(response.json())
-
- return _embed_with_retry(**kwargs)
-
-
-@deprecated(
- since="0.0.29",
- removal="1.0",
- alternative_import="langchain_voyageai.VoyageAIEmbeddings",
-)
-class VoyageEmbeddings(BaseModel, Embeddings):
- """Voyage embedding models.
-
- To use, you should have the environment variable ``VOYAGE_API_KEY`` set with
- your API key or pass it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings import VoyageEmbeddings
-
- voyage = VoyageEmbeddings(voyage_api_key="your-api-key", model="voyage-2")
- text = "This is a test query."
- query_result = voyage.embed_query(text)
- """
-
- model: str
- voyage_api_base: str = "https://api.voyageai.com/v1/embeddings"
- voyage_api_key: Optional[SecretStr] = None
- batch_size: int
- """Maximum number of texts to embed in each API request."""
- max_retries: int = 6
- """Maximum number of retries to make when generating."""
- request_timeout: Optional[Union[float, Tuple[float, float]]] = None
- """Timeout in seconds for the API request."""
- show_progress_bar: bool = False
- """Whether to show a progress bar when embedding. Must have tqdm installed if set
- to True."""
- truncation: bool = True
- """Whether to truncate the input texts to fit within the context length.
-
- If True, over-length input texts will be truncated to fit within the context
- length, before vectorized by the embedding model. If False, an error will be
- raised if any given text exceeds the context length."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and python package exists in environment."""
- values["voyage_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(values, "voyage_api_key", "VOYAGE_API_KEY")
- )
-
- if "model" not in values:
- values["model"] = "voyage-01"
- logger.warning(
- "model will become a required arg for VoyageAIEmbeddings, "
- "we recommend to specify it when using this class. "
- "Currently the default is set to voyage-01."
- )
-
- if "batch_size" not in values:
- values["batch_size"] = (
- 72
- if "model" in values and (values["model"] in ["voyage-2", "voyage-02"])
- else 7
- )
-
- return values
-
- def _invocation_params(
- self, input: List[str], input_type: Optional[str] = None
- ) -> Dict:
- api_key = cast(SecretStr, self.voyage_api_key).get_secret_value()
- params: Dict = {
- "url": self.voyage_api_base,
- "headers": {"Authorization": f"Bearer {api_key}"},
- "json": {
- "model": self.model,
- "input": input,
- "input_type": input_type,
- "truncation": self.truncation,
- },
- "timeout": self.request_timeout,
- }
- return params
-
- def _get_embeddings(
- self,
- texts: List[str],
- batch_size: Optional[int] = None,
- input_type: Optional[str] = None,
- ) -> List[List[float]]:
- embeddings: List[List[float]] = []
-
- if batch_size is None:
- batch_size = self.batch_size
-
- if self.show_progress_bar:
- try:
- from tqdm.auto import tqdm
- except ImportError as e:
- raise ImportError(
- "Must have tqdm installed if `show_progress_bar` is set to True. "
- "Please install with `pip install tqdm`."
- ) from e
-
- _iter = tqdm(range(0, len(texts), batch_size))
- else:
- _iter = range(0, len(texts), batch_size)
-
- if input_type and input_type not in ["query", "document"]:
- raise ValueError(
- f"input_type {input_type} is invalid. Options: None, 'query', "
- "'document'."
- )
-
- for i in _iter:
- response = embed_with_retry(
- self,
- **self._invocation_params(
- input=texts[i : i + batch_size], input_type=input_type
- ),
- )
- embeddings.extend(r["embedding"] for r in response["data"])
-
- return embeddings
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Call out to Voyage Embedding endpoint for embedding search docs.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
- return self._get_embeddings(
- texts, batch_size=self.batch_size, input_type="document"
- )
-
- def embed_query(self, text: str) -> List[float]:
- """Call out to Voyage Embedding endpoint for embedding query text.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embedding for the text.
- """
- return self._get_embeddings(
- [text], batch_size=self.batch_size, input_type="query"
- )[0]
-
- def embed_general_texts(
- self, texts: List[str], *, input_type: Optional[str] = None
- ) -> List[List[float]]:
- """Call out to Voyage Embedding endpoint for embedding general text.
-
- Args:
- texts: The list of texts to embed.
- input_type: Type of the input text. Default to None, meaning the type is
- unspecified. Other options: query, document.
-
- Returns:
- Embedding for the text.
- """
- return self._get_embeddings(
- texts, batch_size=self.batch_size, input_type=input_type
- )
diff --git a/libs/community/langchain_community/embeddings/xinference.py b/libs/community/langchain_community/embeddings/xinference.py
deleted file mode 100644
index 858a2fea41..0000000000
--- a/libs/community/langchain_community/embeddings/xinference.py
+++ /dev/null
@@ -1,139 +0,0 @@
-"""Wrapper around Xinference embedding models."""
-
-from typing import Any, List, Optional
-
-from langchain_core.embeddings import Embeddings
-
-
-class XinferenceEmbeddings(Embeddings):
- """Xinference embedding models.
-
- To use, you should have the xinference library installed:
-
- .. code-block:: bash
-
- pip install xinference
-
- If you're simply using the services provided by Xinference, you can utilize the xinference_client package:
-
- .. code-block:: bash
-
- pip install xinference_client
-
- Check out: https://github.com/xorbitsai/inference
- To run, you need to start a Xinference supervisor on one server and Xinference workers on the other servers.
-
- Example:
- To start a local instance of Xinference, run
-
- .. code-block:: bash
-
- $ xinference
-
- You can also deploy Xinference in a distributed cluster. Here are the steps:
-
- Starting the supervisor:
-
- .. code-block:: bash
-
- $ xinference-supervisor
-
- If you're simply using the services provided by Xinference, you can utilize the xinference_client package:
-
- .. code-block:: bash
-
- pip install xinference_client
-
- Starting the worker:
-
- .. code-block:: bash
-
- $ xinference-worker
-
- Then, launch a model using command line interface (CLI).
-
- Example:
-
- .. code-block:: bash
-
- $ xinference launch -n orca -s 3 -q q4_0
-
- It will return a model UID. Then you can use Xinference Embedding with LangChain.
-
- Example:
-
- .. code-block:: python
-
- from langchain_community.embeddings import XinferenceEmbeddings
-
- xinference = XinferenceEmbeddings(
- server_url="http://0.0.0.0:9997",
- model_uid = {model_uid} # replace model_uid with the model UID return from launching the model
- )
-
- """ # noqa: E501
-
- client: Any
- server_url: Optional[str]
- """URL of the xinference server"""
- model_uid: Optional[str]
- """UID of the launched model"""
-
- def __init__(
- self, server_url: Optional[str] = None, model_uid: Optional[str] = None
- ):
- try:
- from xinference.client import RESTfulClient
- except ImportError:
- try:
- from xinference_client import RESTfulClient
- except ImportError as e:
- raise ImportError(
- "Could not import RESTfulClient from xinference. Please install it"
- " with `pip install xinference` or `pip install xinference_client`."
- ) from e
-
- super().__init__()
-
- if server_url is None:
- raise ValueError("Please provide server URL")
-
- if model_uid is None:
- raise ValueError("Please provide the model UID")
-
- self.server_url = server_url
-
- self.model_uid = model_uid
-
- self.client = RESTfulClient(server_url)
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed a list of documents using Xinference.
- Args:
- texts: The list of texts to embed.
- Returns:
- List of embeddings, one for each text.
- """
-
- model = self.client.get_model(self.model_uid)
-
- embeddings = [
- model.create_embedding(text)["data"][0]["embedding"] for text in texts
- ]
- return [list(map(float, e)) for e in embeddings]
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query of documents using Xinference.
- Args:
- text: The text to embed.
- Returns:
- Embeddings for the text.
- """
-
- model = self.client.get_model(self.model_uid)
-
- embedding_res = model.create_embedding(text)
-
- embedding = embedding_res["data"][0]["embedding"]
-
- return list(map(float, embedding))
diff --git a/libs/community/langchain_community/embeddings/yandex.py b/libs/community/langchain_community/embeddings/yandex.py
deleted file mode 100644
index f7a6ac555c..0000000000
--- a/libs/community/langchain_community/embeddings/yandex.py
+++ /dev/null
@@ -1,212 +0,0 @@
-"""Wrapper around YandexGPT embedding models."""
-
-from __future__ import annotations
-
-import logging
-import time
-from typing import Any, Callable, Dict, List, Sequence
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict, Field, SecretStr
-from tenacity import (
- before_sleep_log,
- retry,
- retry_if_exception_type,
- stop_after_attempt,
- wait_exponential,
-)
-
-logger = logging.getLogger(__name__)
-
-
-class YandexGPTEmbeddings(BaseModel, Embeddings):
- """YandexGPT Embeddings models.
-
- To use, you should have the ``yandexcloud`` python package installed.
-
- There are two authentication options for the service account
- with the ``ai.languageModels.user`` role:
- - You can specify the token in a constructor parameter `iam_token`
- or in an environment variable `YC_IAM_TOKEN`.
- - You can specify the key in a constructor parameter `api_key`
- or in an environment variable `YC_API_KEY`.
-
- To use the default model specify the folder ID in a parameter `folder_id`
- or in an environment variable `YC_FOLDER_ID`.
-
- Example:
- .. code-block:: python
-
- from langchain_community.embeddings.yandex import YandexGPTEmbeddings
- embeddings = YandexGPTEmbeddings(iam_token="t1.9eu...", folder_id=)
- """ # noqa: E501
-
- iam_token: SecretStr = "" # type: ignore[assignment]
- """Yandex Cloud IAM token for service account
- with the `ai.languageModels.user` role"""
- api_key: SecretStr = "" # type: ignore[assignment]
- """Yandex Cloud Api Key for service account
- with the `ai.languageModels.user` role"""
- model_uri: str = Field(default="", alias="query_model_uri")
- """Query model uri to use."""
- doc_model_uri: str = ""
- """Doc model uri to use."""
- folder_id: str = ""
- """Yandex Cloud folder ID"""
- doc_model_name: str = "text-search-doc"
- """Doc model name to use."""
- model_name: str = Field(default="text-search-query", alias="query_model_name")
- """Query model name to use."""
- model_version: str = "latest"
- """Model version to use."""
- url: str = "llm.api.cloud.yandex.net:443"
- """The url of the API."""
- max_retries: int = 6
- """Maximum number of retries to make when generating."""
- sleep_interval: float = 0.0
- """Delay between API requests"""
- disable_request_logging: bool = False
- """YandexGPT API logs all request data by default.
- If you provide personal data, confidential information, disable logging."""
- grpc_metadata: Sequence
-
- model_config = ConfigDict(populate_by_name=True, protected_namespaces=())
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that iam token exists in environment."""
-
- iam_token = convert_to_secret_str(
- get_from_dict_or_env(values, "iam_token", "YC_IAM_TOKEN", "")
- )
- values["iam_token"] = iam_token
- api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "api_key", "YC_API_KEY", "")
- )
- values["api_key"] = api_key
- folder_id = get_from_dict_or_env(values, "folder_id", "YC_FOLDER_ID", "")
- values["folder_id"] = folder_id
- if api_key.get_secret_value() == "" and iam_token.get_secret_value() == "":
- raise ValueError("Either 'YC_API_KEY' or 'YC_IAM_TOKEN' must be provided.")
- if values["iam_token"]:
- values["grpc_metadata"] = [
- ("authorization", f"Bearer {values['iam_token'].get_secret_value()}")
- ]
- if values["folder_id"]:
- values["grpc_metadata"].append(("x-folder-id", values["folder_id"]))
- else:
- values["grpc_metadata"] = [
- ("authorization", f"Api-Key {values['api_key'].get_secret_value()}"),
- ]
-
- if not values.get("doc_model_uri"):
- if values["folder_id"] == "":
- raise ValueError("'doc_model_uri' or 'folder_id' must be provided.")
- values["doc_model_uri"] = (
- f"emb://{values['folder_id']}/{values['doc_model_name']}/{values['model_version']}"
- )
- if not values.get("model_uri"):
- if values["folder_id"] == "":
- raise ValueError("'model_uri' or 'folder_id' must be provided.")
- values["model_uri"] = (
- f"emb://{values['folder_id']}/{values['model_name']}/{values['model_version']}"
- )
- if values["disable_request_logging"]:
- values["grpc_metadata"].append(
- (
- "x-data-logging-enabled",
- "false",
- )
- )
- return values
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """Embed documents using a YandexGPT embeddings models.
-
- Args:
- texts: The list of texts to embed.
-
- Returns:
- List of embeddings, one for each text.
- """
-
- return _embed_with_retry(self, texts=texts)
-
- def embed_query(self, text: str) -> List[float]:
- """Embed a query using a YandexGPT embeddings models.
-
- Args:
- text: The text to embed.
-
- Returns:
- Embeddings for the text.
- """
- return _embed_with_retry(self, texts=[text], embed_query=True)[0]
-
-
-def _create_retry_decorator(llm: YandexGPTEmbeddings) -> Callable[[Any], Any]:
- from grpc import RpcError
-
- min_seconds = 1
- max_seconds = 60
- return retry(
- reraise=True,
- stop=stop_after_attempt(llm.max_retries),
- wait=wait_exponential(multiplier=1, min=min_seconds, max=max_seconds),
- retry=(retry_if_exception_type((RpcError))),
- before_sleep=before_sleep_log(logger, logging.WARNING),
- )
-
-
-def _embed_with_retry(llm: YandexGPTEmbeddings, **kwargs: Any) -> Any:
- """Use tenacity to retry the embedding call."""
- retry_decorator = _create_retry_decorator(llm)
-
- @retry_decorator
- def _completion_with_retry(**_kwargs: Any) -> Any:
- return _make_request(llm, **_kwargs)
-
- return _completion_with_retry(**kwargs)
-
-
-def _make_request(self: YandexGPTEmbeddings, texts: List[str], **kwargs): # type: ignore[no-untyped-def]
- try:
- import grpc
-
- try:
- from yandex.cloud.ai.foundation_models.v1.embedding.embedding_service_pb2 import ( # noqa: E501
- TextEmbeddingRequest,
- )
- from yandex.cloud.ai.foundation_models.v1.embedding.embedding_service_pb2_grpc import ( # noqa: E501
- EmbeddingsServiceStub,
- )
- except ModuleNotFoundError:
- from yandex.cloud.ai.foundation_models.v1.foundation_models_service_pb2 import ( # noqa: E501
- TextEmbeddingRequest,
- )
- from yandex.cloud.ai.foundation_models.v1.foundation_models_service_pb2_grpc import ( # noqa: E501
- EmbeddingsServiceStub,
- )
- except ImportError as e:
- raise ImportError(
- "Please install YandexCloud SDK with `pip install yandexcloud` \
- or upgrade it to recent version."
- ) from e
- result = []
- channel_credentials = grpc.ssl_channel_credentials()
- channel = grpc.secure_channel(self.url, channel_credentials)
- # Use the query model if embed_query is True
- if kwargs.get("embed_query"):
- model_uri = self.model_uri
- else:
- model_uri = self.doc_model_uri
-
- for text in texts:
- request = TextEmbeddingRequest(model_uri=model_uri, text=text)
- stub = EmbeddingsServiceStub(channel)
- res = stub.TextEmbedding(request, metadata=self.grpc_metadata)
- result.append(list(res.embedding))
- time.sleep(self.sleep_interval)
-
- return result
diff --git a/libs/community/langchain_community/embeddings/zhipuai.py b/libs/community/langchain_community/embeddings/zhipuai.py
deleted file mode 100644
index 73ced5fa01..0000000000
--- a/libs/community/langchain_community/embeddings/zhipuai.py
+++ /dev/null
@@ -1,128 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-from langchain_core.embeddings import Embeddings
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import BaseModel, Field, model_validator
-
-
-class ZhipuAIEmbeddings(BaseModel, Embeddings):
- """ZhipuAI embedding model integration.
-
- Setup:
-
- To use, you should have the ``zhipuai`` python package installed, and the
- environment variable ``ZHIPU_API_KEY`` set with your API KEY.
-
- More instructions about ZhipuAi Embeddings, you can get it
- from https://open.bigmodel.cn/dev/api#vector
-
- .. code-block:: bash
-
- pip install -U zhipuai
- export ZHIPU_API_KEY="your-api-key"
-
- Key init args — completion params:
- model: Optional[str]
- Name of ZhipuAI model to use.
- api_key: str
- Automatically inferred from env var `ZHIPU_API_KEY` if not provided.
-
- See full list of supported init args and their descriptions in the params section.
-
- Instantiate:
-
- .. code-block:: python
-
- from langchain_community.embeddings import ZhipuAIEmbeddings
-
- embed = ZhipuAIEmbeddings(
- model="embedding-2",
- # api_key="...",
- )
-
- Embed single text:
- .. code-block:: python
-
- input_text = "The meaning of life is 42"
- embed.embed_query(input_text)
-
- .. code-block:: python
-
- [-0.003832892, 0.049372625, -0.035413884, -0.019301128, 0.0068899863, 0.01248398, -0.022153955, 0.006623926, 0.00778216, 0.009558191, ...]
-
-
- Embed multiple text:
- .. code-block:: python
-
- input_texts = ["This is a test query1.", "This is a test query2."]
- embed.embed_documents(input_texts)
-
- .. code-block:: python
-
- [
- [0.0083934665, 0.037985895, -0.06684559, -0.039616987, 0.015481004, -0.023952313, ...],
- [-0.02713102, -0.005470169, 0.032321047, 0.042484466, 0.023290444, 0.02170547, ...]
- ]
- """ # noqa: E501
-
- client: Any = Field(default=None, exclude=True) #: :meta private:
- model: str = Field(default="embedding-2")
- """Model name"""
- api_key: str
- """Automatically inferred from env var `ZHIPU_API_KEY` if not provided."""
- dimensions: Optional[int] = None
- """The number of dimensions the resulting output embeddings should have.
-
- Only supported in `embedding-3` and later models.
- """
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that auth token exists in environment."""
- values["api_key"] = get_from_dict_or_env(values, "api_key", "ZHIPUAI_API_KEY")
- try:
- from zhipuai import ZhipuAI
-
- values["client"] = ZhipuAI(api_key=values["api_key"])
- except ImportError:
- raise ImportError(
- "Could not import zhipuai python package."
- "Please install it with `pip install zhipuai`."
- )
- return values
-
- def embed_query(self, text: str) -> List[float]:
- """
- Embeds a text using the AutoVOT algorithm.
-
- Args:
- text: A text to embed.
-
- Returns:
- Input document's embedded list.
- """
- resp = self.embed_documents([text])
- return resp[0]
-
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- """
- Embeds a list of text documents using the AutoVOT algorithm.
-
- Args:
- texts: A list of text documents to embed.
-
- Returns:
- A list of embeddings for each document in the input list.
- Each embedding is represented as a list of float values.
- """
- if self.dimensions is not None:
- resp = self.client.embeddings.create(
- model=self.model,
- input=texts,
- dimensions=self.dimensions,
- )
- else:
- resp = self.client.embeddings.create(model=self.model, input=texts)
- embeddings = [r.embedding for r in resp.data]
- return embeddings
diff --git a/libs/community/langchain_community/example_selectors/__init__.py b/libs/community/langchain_community/example_selectors/__init__.py
deleted file mode 100644
index d29bf73723..0000000000
--- a/libs/community/langchain_community/example_selectors/__init__.py
+++ /dev/null
@@ -1,18 +0,0 @@
-"""**Example selector** implements logic for selecting examples to include them
-in prompts.
-This allows us to select examples that are most relevant to the input.
-
-There could be multiple strategies for selecting examples. For example, one could
-select examples based on the similarity of the input to the examples. Another
-strategy could be to select examples based on the diversity of the examples.
-"""
-
-from langchain_community.example_selectors.ngram_overlap import (
- NGramOverlapExampleSelector,
- ngram_overlap_score,
-)
-
-__all__ = [
- "NGramOverlapExampleSelector",
- "ngram_overlap_score",
-]
diff --git a/libs/community/langchain_community/example_selectors/ngram_overlap.py b/libs/community/langchain_community/example_selectors/ngram_overlap.py
deleted file mode 100644
index 92577acd56..0000000000
--- a/libs/community/langchain_community/example_selectors/ngram_overlap.py
+++ /dev/null
@@ -1,116 +0,0 @@
-"""Select and order examples based on ngram overlap score (sentence_bleu score).
-
-https://www.nltk.org/_modules/nltk/translate/bleu_score.html
-https://aclanthology.org/P02-1040.pdf
-"""
-
-from typing import Any, Dict, List
-
-import numpy as np
-from langchain_core.example_selectors import BaseExampleSelector
-from langchain_core.prompts import PromptTemplate
-from pydantic import BaseModel, model_validator
-
-
-def ngram_overlap_score(source: List[str], example: List[str]) -> float:
- """Compute ngram overlap score of source and example as sentence_bleu score
- from NLTK package.
-
- Use sentence_bleu with method1 smoothing function and auto reweighting.
- Return float value between 0.0 and 1.0 inclusive.
- https://www.nltk.org/_modules/nltk/translate/bleu_score.html
- https://aclanthology.org/P02-1040.pdf
- """
- from nltk.translate.bleu_score import (
- SmoothingFunction,
- sentence_bleu,
- )
-
- hypotheses = source[0].split()
- references = [s.split() for s in example]
-
- return float(
- sentence_bleu(
- references,
- hypotheses,
- smoothing_function=SmoothingFunction().method1,
- auto_reweigh=True,
- )
- )
-
-
-class NGramOverlapExampleSelector(BaseExampleSelector, BaseModel):
- """Select and order examples based on ngram overlap score (sentence_bleu score
- from NLTK package).
-
- https://www.nltk.org/_modules/nltk/translate/bleu_score.html
- https://aclanthology.org/P02-1040.pdf
- """
-
- examples: List[dict]
- """A list of the examples that the prompt template expects."""
-
- example_prompt: PromptTemplate
- """Prompt template used to format the examples."""
-
- threshold: float = -1.0
- """Threshold at which algorithm stops. Set to -1.0 by default.
-
- For negative threshold:
- select_examples sorts examples by ngram_overlap_score, but excludes none.
- For threshold greater than 1.0:
- select_examples excludes all examples, and returns an empty list.
- For threshold equal to 0.0:
- select_examples sorts examples by ngram_overlap_score,
- and excludes examples with no ngram overlap with input.
- """
-
- @model_validator(mode="before")
- @classmethod
- def check_dependencies(cls, values: Dict) -> Any:
- """Check that valid dependencies exist."""
- try:
- from nltk.translate.bleu_score import ( # noqa: F401
- SmoothingFunction,
- sentence_bleu,
- )
- except ImportError as e:
- raise ImportError(
- "Not all the correct dependencies for this ExampleSelect exist."
- "Please install nltk with `pip install nltk`."
- ) from e
-
- return values
-
- def add_example(self, example: Dict[str, str]) -> None:
- """Add new example to list."""
- self.examples.append(example)
-
- def select_examples(self, input_variables: Dict[str, str]) -> List[dict]:
- """Return list of examples sorted by ngram_overlap_score with input.
-
- Descending order.
- Excludes any examples with ngram_overlap_score less than or equal to threshold.
- """
- inputs = list(input_variables.values())
- examples = []
- k = len(self.examples)
- score = [0.0] * k
- first_prompt_template_key = self.example_prompt.input_variables[0]
-
- for i in range(k):
- score[i] = ngram_overlap_score(
- inputs, [self.examples[i][first_prompt_template_key]]
- )
-
- while True:
- arg_max = np.argmax(score)
- if (score[arg_max] < self.threshold) or abs(
- score[arg_max] - self.threshold
- ) < 1e-9:
- break
-
- examples.append(self.examples[arg_max])
- score[arg_max] = self.threshold - 1.0
-
- return examples
diff --git a/libs/community/langchain_community/graph_vectorstores/__init__.py b/libs/community/langchain_community/graph_vectorstores/__init__.py
deleted file mode 100644
index b01e20a2d6..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/__init__.py
+++ /dev/null
@@ -1,157 +0,0 @@
-""".. title:: Graph Vector Store
-
-Graph Vector Store
-==================
-
-Sometimes embedding models don't capture all the important relationships between
-documents.
-Graph Vector Stores are an extension to both vector stores and retrievers that allow
-documents to be explicitly connected to each other.
-
-Graph vector store retrievers use both vector similarity and links to find documents
-related to an unstructured query.
-
-Graphs allow linking between documents.
-Each document identifies tags that link to and from it.
-For example, a paragraph of text may be linked to URLs based on the anchor tags in
-it's content and linked from the URL(s) it is published at.
-
-`Link extractors `
-can be used to extract links from documents.
-
-Example::
-
- graph_vector_store = CassandraGraphVectorStore()
- link_extractor = HtmlLinkExtractor()
- links = link_extractor.extract_one(HtmlInput(document.page_content, "http://mysite"))
- add_links(document, links)
- graph_vector_store.add_document(document)
-
-.. seealso::
-
- - :class:`How to use a graph vector store as a retriever `
- - :class:`How to create links between documents `
- - :class:`How to link Documents on hyperlinks in HTML `
- - :class:`How to link Documents on common keywords (using KeyBERT) `
- - :class:`How to link Documents on common named entities (using GliNER) `
- - `langchain-jieba: link extraction tailored for Chinese language `_
-
-Get started
------------
-
-We chunk the State of the Union text and split it into documents::
-
- from langchain_community.document_loaders import TextLoader
- from langchain_text_splitters import CharacterTextSplitter
-
- raw_documents = TextLoader("state_of_the_union.txt").load()
- text_splitter = CharacterTextSplitter(chunk_size=1000, chunk_overlap=0)
- documents = text_splitter.split_documents(raw_documents)
-
-Links can be added to documents manually but it's easier to use a
-:class:`~langchain_community.graph_vectorstores.extractors.link_extractor.LinkExtractor`.
-Several common link extractors are available and you can build your own.
-For this guide, we'll use the
-:class:`~langchain_community.graph_vectorstores.extractors.keybert_link_extractor.KeybertLinkExtractor`
-which uses the KeyBERT model to tag documents with keywords and uses these keywords to
-create links between documents::
-
- from langchain_community.graph_vectorstores.extractors import KeybertLinkExtractor
- from langchain_community.graph_vectorstores.links import add_links
-
- extractor = KeybertLinkExtractor()
-
- for doc in documents:
- add_links(doc, extractor.extract_one(doc))
-
-Create the graph vector store and add documents
------------------------------------------------
-
-We'll use an Apache Cassandra or Astra DB database as an example.
-We create a
-:class:`~langchain_community.graph_vectorstores.cassandra.CassandraGraphVectorStore`
-from the documents and an :class:`~langchain_openai.embeddings.base.OpenAIEmbeddings`
-model::
-
- import cassio
- from langchain_community.graph_vectorstores import CassandraGraphVectorStore
- from langchain_openai import OpenAIEmbeddings
-
- # Initialize cassio and the Cassandra session from the environment variables
- cassio.init(auto=True)
-
- store = CassandraGraphVectorStore.from_documents(
- embedding=OpenAIEmbeddings(),
- documents=documents,
- )
-
-
-Similarity search
------------------
-
-If we don't traverse the graph, a graph vector store behaves like a regular vector
-store.
-So all methods available in a vector store are also available in a graph vector store.
-The :meth:`~langchain_community.graph_vectorstores.base.GraphVectorStore.similarity_search`
-method returns documents similar to a query without considering
-the links between documents::
-
- docs = store.similarity_search(
- "What did the president say about Ketanji Brown Jackson?"
- )
-
-Traversal search
-----------------
-
-The :meth:`~langchain_community.graph_vectorstores.base.GraphVectorStore.traversal_search`
-method returns documents similar to a query considering the links
-between documents. It first does a similarity search and then traverses the graph to
-find linked documents::
-
- docs = list(
- store.traversal_search("What did the president say about Ketanji Brown Jackson?")
- )
-
-Async methods
--------------
-
-The graph vector store has async versions of the methods prefixed with ``a``::
-
- docs = [
- doc
- async for doc in store.atraversal_search(
- "What did the president say about Ketanji Brown Jackson?"
- )
- ]
-
-Graph vector store retriever
-----------------------------
-
-The graph vector store can be converted to a retriever.
-It is similar to the vector store retriever but it also has traversal search methods
-such as ``traversal`` and ``mmr_traversal``::
-
- retriever = store.as_retriever(search_type="mmr_traversal")
- docs = retriever.invoke("What did the president say about Ketanji Brown Jackson?")
-
-""" # noqa: E501
-
-from langchain_community.graph_vectorstores.base import (
- GraphVectorStore,
- GraphVectorStoreRetriever,
- Node,
-)
-from langchain_community.graph_vectorstores.cassandra import CassandraGraphVectorStore
-from langchain_community.graph_vectorstores.links import (
- Link,
-)
-from langchain_community.graph_vectorstores.mmr_helper import MmrHelper
-
-__all__ = [
- "GraphVectorStore",
- "GraphVectorStoreRetriever",
- "Node",
- "Link",
- "CassandraGraphVectorStore",
- "MmrHelper",
-]
diff --git a/libs/community/langchain_community/graph_vectorstores/base.py b/libs/community/langchain_community/graph_vectorstores/base.py
deleted file mode 100644
index 58be972b14..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/base.py
+++ /dev/null
@@ -1,917 +0,0 @@
-from __future__ import annotations
-
-import logging
-from abc import abstractmethod
-from collections.abc import AsyncIterable, Collection, Iterable, Iterator
-from typing import (
- Any,
- ClassVar,
- Optional,
- Sequence,
- cast,
-)
-
-from langchain_core._api import deprecated
-from langchain_core.callbacks import (
- AsyncCallbackManagerForRetrieverRun,
- CallbackManagerForRetrieverRun,
-)
-from langchain_core.documents import Document
-from langchain_core.load import Serializable
-from langchain_core.runnables import run_in_executor
-from langchain_core.vectorstores import VectorStore, VectorStoreRetriever
-from pydantic import Field
-
-from langchain_community.graph_vectorstores.links import METADATA_LINKS_KEY, Link
-
-logger = logging.getLogger(__name__)
-
-
-def _has_next(iterator: Iterator) -> bool:
- """Checks if the iterator has more elements.
- Warning: consumes an element from the iterator"""
- sentinel = object()
- return next(iterator, sentinel) is not sentinel
-
-
-DEPRECATION_ADDENDUM = (
- "See https://datastax.github.io/graph-rag/guide/migration/"
- "#from-langchain-graphvectorstore for migration instructions."
-)
-
-
-@deprecated(
- since="0.3.21",
- removal="0.5",
- addendum=DEPRECATION_ADDENDUM,
-)
-class Node(Serializable):
- """Node in the GraphVectorStore.
-
- Edges exist from nodes with an outgoing link to nodes with a matching incoming link.
-
- For instance two nodes `a` and `b` connected over a hyperlink ``https://some-url``
- would look like:
-
- .. code-block:: python
-
- [
- Node(
- id="a",
- text="some text a",
- links= [
- Link(kind="hyperlink", tag="https://some-url", direction="incoming")
- ],
- ),
- Node(
- id="b",
- text="some text b",
- links= [
- Link(kind="hyperlink", tag="https://some-url", direction="outgoing")
- ],
- )
- ]
- """
-
- id: Optional[str] = None
- """Unique ID for the node. Will be generated by the GraphVectorStore if not set."""
- text: str
- """Text contained by the node."""
- metadata: dict = Field(default_factory=dict)
- """Metadata for the node."""
- links: list[Link] = Field(default_factory=list)
- """Links associated with the node."""
-
-
-def _texts_to_nodes(
- texts: Iterable[str],
- metadatas: Optional[Iterable[dict]],
- ids: Optional[Iterable[str]],
-) -> Iterator[Node]:
- metadatas_it = iter(metadatas) if metadatas else None
- ids_it = iter(ids) if ids else None
- for text in texts:
- try:
- _metadata = next(metadatas_it).copy() if metadatas_it else {}
- except StopIteration as e:
- raise ValueError("texts iterable longer than metadatas") from e
- try:
- _id = next(ids_it) if ids_it else None
- except StopIteration as e:
- raise ValueError("texts iterable longer than ids") from e
-
- links = _metadata.pop(METADATA_LINKS_KEY, [])
- if not isinstance(links, list):
- links = list(links)
- yield Node(
- id=_id,
- metadata=_metadata,
- text=text,
- links=links,
- )
- if ids_it and _has_next(ids_it):
- raise ValueError("ids iterable longer than texts")
- if metadatas_it and _has_next(metadatas_it):
- raise ValueError("metadatas iterable longer than texts")
-
-
-def _documents_to_nodes(documents: Iterable[Document]) -> Iterator[Node]:
- for doc in documents:
- metadata = doc.metadata.copy()
- links = metadata.pop(METADATA_LINKS_KEY, [])
- if not isinstance(links, list):
- links = list(links)
- yield Node(
- id=doc.id,
- metadata=metadata,
- text=doc.page_content,
- links=links,
- )
-
-
-@deprecated(
- since="0.3.21",
- removal="0.5",
- addendum=DEPRECATION_ADDENDUM,
-)
-def nodes_to_documents(nodes: Iterable[Node]) -> Iterator[Document]:
- """Convert nodes to documents.
-
- Args:
- nodes: The nodes to convert to documents.
- Returns:
- The documents generated from the nodes.
- """
- for node in nodes:
- metadata = node.metadata.copy()
- metadata[METADATA_LINKS_KEY] = [
- # Convert the core `Link` (from the node) back to the local `Link`.
- Link(kind=link.kind, direction=link.direction, tag=link.tag)
- for link in node.links
- ]
-
- yield Document(
- id=node.id,
- page_content=node.text,
- metadata=metadata,
- )
-
-
-@deprecated(
- since="0.3.21",
- removal="0.5",
- addendum=DEPRECATION_ADDENDUM,
-)
-class GraphVectorStore(VectorStore):
- """A hybrid vector-and-graph graph store.
-
- Document chunks support vector-similarity search as well as edges linking
- chunks based on structural and semantic properties.
-
- .. versionadded:: 0.3.1
- """
-
- @abstractmethod
- def add_nodes(
- self,
- nodes: Iterable[Node],
- **kwargs: Any,
- ) -> Iterable[str]:
- """Add nodes to the graph store.
-
- Args:
- nodes: the nodes to add.
- **kwargs: Additional keyword arguments.
- """
-
- async def aadd_nodes(
- self,
- nodes: Iterable[Node],
- **kwargs: Any,
- ) -> AsyncIterable[str]:
- """Add nodes to the graph store.
-
- Args:
- nodes: the nodes to add.
- **kwargs: Additional keyword arguments.
- """
- iterator = iter(await run_in_executor(None, self.add_nodes, nodes, **kwargs))
- done = object()
- while True:
- doc = await run_in_executor(None, next, iterator, done)
- if doc is done:
- break
- yield doc # type: ignore[misc]
-
- def add_texts(
- self,
- texts: Iterable[str],
- metadatas: Optional[Iterable[dict]] = None,
- *,
- ids: Optional[Iterable[str]] = None,
- **kwargs: Any,
- ) -> list[str]:
- """Run more texts through the embeddings and add to the vector store.
-
- The Links present in the metadata field `links` will be extracted to create
- the `Node` links.
-
- Eg if nodes `a` and `b` are connected over a hyperlink `https://some-url`, the
- function call would look like:
-
- .. code-block:: python
-
- store.add_texts(
- ids=["a", "b"],
- texts=["some text a", "some text b"],
- metadatas=[
- {
- "links": [
- Link.incoming(kind="hyperlink", tag="https://some-url")
- ]
- },
- {
- "links": [
- Link.outgoing(kind="hyperlink", tag="https://some-url")
- ]
- },
- ],
- )
-
- Args:
- texts: Iterable of strings to add to the vector store.
- metadatas: Optional list of metadatas associated with the texts.
- The metadata key `links` shall be an iterable of
- :py:class:`~langchain_community.graph_vectorstores.links.Link`.
- ids: Optional list of IDs associated with the texts.
- **kwargs: vector store specific parameters.
-
- Returns:
- List of ids from adding the texts into the vector store.
- """
- nodes = _texts_to_nodes(texts, metadatas, ids)
- return list(self.add_nodes(nodes, **kwargs))
-
- async def aadd_texts(
- self,
- texts: Iterable[str],
- metadatas: Optional[Iterable[dict]] = None,
- *,
- ids: Optional[Iterable[str]] = None,
- **kwargs: Any,
- ) -> list[str]:
- """Run more texts through the embeddings and add to the vector store.
-
- The Links present in the metadata field `links` will be extracted to create
- the `Node` links.
-
- Eg if nodes `a` and `b` are connected over a hyperlink `https://some-url`, the
- function call would look like:
-
- .. code-block:: python
-
- await store.aadd_texts(
- ids=["a", "b"],
- texts=["some text a", "some text b"],
- metadatas=[
- {
- "links": [
- Link.incoming(kind="hyperlink", tag="https://some-url")
- ]
- },
- {
- "links": [
- Link.outgoing(kind="hyperlink", tag="https://some-url")
- ]
- },
- ],
- )
-
- Args:
- texts: Iterable of strings to add to the vector store.
- metadatas: Optional list of metadatas associated with the texts.
- The metadata key `links` shall be an iterable of
- :py:class:`~langchain_community.graph_vectorstores.links.Link`.
- ids: Optional list of IDs associated with the texts.
- **kwargs: vector store specific parameters.
-
- Returns:
- List of ids from adding the texts into the vector store.
- """
- nodes = _texts_to_nodes(texts, metadatas, ids)
- return [_id async for _id in self.aadd_nodes(nodes, **kwargs)]
-
- def add_documents(
- self,
- documents: Iterable[Document],
- **kwargs: Any,
- ) -> list[str]:
- """Run more documents through the embeddings and add to the vector store.
-
- The Links present in the document metadata field `links` will be extracted to
- create the `Node` links.
-
- Eg if nodes `a` and `b` are connected over a hyperlink `https://some-url`, the
- function call would look like:
-
- .. code-block:: python
-
- store.add_documents(
- [
- Document(
- id="a",
- page_content="some text a",
- metadata={
- "links": [
- Link.incoming(kind="hyperlink", tag="http://some-url")
- ]
- }
- ),
- Document(
- id="b",
- page_content="some text b",
- metadata={
- "links": [
- Link.outgoing(kind="hyperlink", tag="http://some-url")
- ]
- }
- ),
- ]
-
- )
-
- Args:
- documents: Documents to add to the vector store.
- The document's metadata key `links` shall be an iterable of
- :py:class:`~langchain_community.graph_vectorstores.links.Link`.
-
- Returns:
- List of IDs of the added texts.
- """
- nodes = _documents_to_nodes(documents)
- return list(self.add_nodes(nodes, **kwargs))
-
- async def aadd_documents(
- self,
- documents: Iterable[Document],
- **kwargs: Any,
- ) -> list[str]:
- """Run more documents through the embeddings and add to the vector store.
-
- The Links present in the document metadata field `links` will be extracted to
- create the `Node` links.
-
- Eg if nodes `a` and `b` are connected over a hyperlink `https://some-url`, the
- function call would look like:
-
- .. code-block:: python
-
- store.add_documents(
- [
- Document(
- id="a",
- page_content="some text a",
- metadata={
- "links": [
- Link.incoming(kind="hyperlink", tag="http://some-url")
- ]
- }
- ),
- Document(
- id="b",
- page_content="some text b",
- metadata={
- "links": [
- Link.outgoing(kind="hyperlink", tag="http://some-url")
- ]
- }
- ),
- ]
-
- )
-
- Args:
- documents: Documents to add to the vector store.
- The document's metadata key `links` shall be an iterable of
- :py:class:`~langchain_community.graph_vectorstores.links.Link`.
-
- Returns:
- List of IDs of the added texts.
- """
- nodes = _documents_to_nodes(documents)
- return [_id async for _id in self.aadd_nodes(nodes, **kwargs)]
-
- @abstractmethod
- def traversal_search(
- self,
- query: str,
- *,
- k: int = 4,
- depth: int = 1,
- filter: dict[str, Any] | None = None, # noqa: A002
- **kwargs: Any,
- ) -> Iterable[Document]:
- """Retrieve documents from traversing this graph store.
-
- First, `k` nodes are retrieved using a search for each `query` string.
- Then, additional nodes are discovered up to the given `depth` from those
- starting nodes.
-
- Args:
- query: The query string.
- k: The number of Documents to return from the initial search.
- Defaults to 4. Applies to each of the query strings.
- depth: The maximum depth of edges to traverse. Defaults to 1.
- filter: Optional metadata to filter the results.
- **kwargs: Additional keyword arguments.
- Returns:
- Collection of retrieved documents.
- """
-
- async def atraversal_search(
- self,
- query: str,
- *,
- k: int = 4,
- depth: int = 1,
- filter: dict[str, Any] | None = None, # noqa: A002
- **kwargs: Any,
- ) -> AsyncIterable[Document]:
- """Retrieve documents from traversing this graph store.
-
- First, `k` nodes are retrieved using a search for each `query` string.
- Then, additional nodes are discovered up to the given `depth` from those
- starting nodes.
-
- Args:
- query: The query string.
- k: The number of Documents to return from the initial search.
- Defaults to 4. Applies to each of the query strings.
- depth: The maximum depth of edges to traverse. Defaults to 1.
- filter: Optional metadata to filter the results.
- **kwargs: Additional keyword arguments.
- Returns:
- Collection of retrieved documents.
- """
- iterator = iter(
- await run_in_executor(
- None,
- self.traversal_search,
- query,
- k=k,
- depth=depth,
- filter=filter,
- **kwargs,
- )
- )
- done = object()
- while True:
- doc = await run_in_executor(None, next, iterator, done)
- if doc is done:
- break
- yield doc # type: ignore[misc]
-
- @abstractmethod
- def mmr_traversal_search(
- self,
- query: str,
- *,
- initial_roots: Sequence[str] = (),
- k: int = 4,
- depth: int = 2,
- fetch_k: int = 100,
- adjacent_k: int = 10,
- lambda_mult: float = 0.5,
- score_threshold: float = float("-inf"),
- filter: dict[str, Any] | None = None, # noqa: A002
- **kwargs: Any,
- ) -> Iterable[Document]:
- """Retrieve documents from this graph store using MMR-traversal.
-
- This strategy first retrieves the top `fetch_k` results by similarity to
- the question. It then selects the top `k` results based on
- maximum-marginal relevance using the given `lambda_mult`.
-
- At each step, it considers the (remaining) documents from `fetch_k` as
- well as any documents connected by edges to a selected document
- retrieved based on similarity (a "root").
-
- Args:
- query: The query string to search for.
- initial_roots: Optional list of document IDs to use for initializing search.
- The top `adjacent_k` nodes adjacent to each initial root will be
- included in the set of initial candidates. To fetch only in the
- neighborhood of these nodes, set `fetch_k = 0`.
- k: Number of Documents to return. Defaults to 4.
- fetch_k: Number of Documents to fetch via similarity.
- Defaults to 100.
- adjacent_k: Number of adjacent Documents to fetch.
- Defaults to 10.
- depth: Maximum depth of a node (number of edges) from a node
- retrieved via similarity. Defaults to 2.
- lambda_mult: Number between 0 and 1 that determines the degree
- of diversity among the results with 0 corresponding to maximum
- diversity and 1 to minimum diversity. Defaults to 0.5.
- score_threshold: Only documents with a score greater than or equal
- this threshold will be chosen. Defaults to negative infinity.
- filter: Optional metadata to filter the results.
- **kwargs: Additional keyword arguments.
- """
-
- async def ammr_traversal_search(
- self,
- query: str,
- *,
- initial_roots: Sequence[str] = (),
- k: int = 4,
- depth: int = 2,
- fetch_k: int = 100,
- adjacent_k: int = 10,
- lambda_mult: float = 0.5,
- score_threshold: float = float("-inf"),
- filter: dict[str, Any] | None = None, # noqa: A002
- **kwargs: Any,
- ) -> AsyncIterable[Document]:
- """Retrieve documents from this graph store using MMR-traversal.
-
- This strategy first retrieves the top `fetch_k` results by similarity to
- the question. It then selects the top `k` results based on
- maximum-marginal relevance using the given `lambda_mult`.
-
- At each step, it considers the (remaining) documents from `fetch_k` as
- well as any documents connected by edges to a selected document
- retrieved based on similarity (a "root").
-
- Args:
- query: The query string to search for.
- initial_roots: Optional list of document IDs to use for initializing search.
- The top `adjacent_k` nodes adjacent to each initial root will be
- included in the set of initial candidates. To fetch only in the
- neighborhood of these nodes, set `fetch_k = 0`.
- k: Number of Documents to return. Defaults to 4.
- fetch_k: Number of Documents to fetch via similarity.
- Defaults to 100.
- adjacent_k: Number of adjacent Documents to fetch.
- Defaults to 10.
- depth: Maximum depth of a node (number of edges) from a node
- retrieved via similarity. Defaults to 2.
- lambda_mult: Number between 0 and 1 that determines the degree
- of diversity among the results with 0 corresponding to maximum
- diversity and 1 to minimum diversity. Defaults to 0.5.
- score_threshold: Only documents with a score greater than or equal
- this threshold will be chosen. Defaults to negative infinity.
- filter: Optional metadata to filter the results.
- **kwargs: Additional keyword arguments.
- """
- iterator = iter(
- await run_in_executor(
- None,
- self.mmr_traversal_search,
- query,
- initial_roots=initial_roots,
- k=k,
- fetch_k=fetch_k,
- adjacent_k=adjacent_k,
- depth=depth,
- lambda_mult=lambda_mult,
- score_threshold=score_threshold,
- filter=filter,
- **kwargs,
- )
- )
- done = object()
- while True:
- doc = await run_in_executor(None, next, iterator, done)
- if doc is done:
- break
- yield doc # type: ignore[misc]
-
- def similarity_search(
- self, query: str, k: int = 4, **kwargs: Any
- ) -> list[Document]:
- return list(self.traversal_search(query, k=k, depth=0))
-
- def max_marginal_relevance_search(
- self,
- query: str,
- k: int = 4,
- fetch_k: int = 20,
- lambda_mult: float = 0.5,
- **kwargs: Any,
- ) -> list[Document]:
- if kwargs.get("depth", 0) > 0:
- logger.warning(
- "'mmr' search started with depth > 0. "
- "Maybe you meant to do a 'mmr_traversal' search?"
- )
- return list(
- self.mmr_traversal_search(
- query, k=k, fetch_k=fetch_k, lambda_mult=lambda_mult, depth=0
- )
- )
-
- async def asimilarity_search(
- self, query: str, k: int = 4, **kwargs: Any
- ) -> list[Document]:
- return [doc async for doc in self.atraversal_search(query, k=k, depth=0)]
-
- def search(self, query: str, search_type: str, **kwargs: Any) -> list[Document]:
- if search_type == "similarity":
- return self.similarity_search(query, **kwargs)
- elif search_type == "similarity_score_threshold":
- docs_and_similarities = self.similarity_search_with_relevance_scores(
- query, **kwargs
- )
- return [doc for doc, _ in docs_and_similarities]
- elif search_type == "mmr":
- return self.max_marginal_relevance_search(query, **kwargs)
- elif search_type == "traversal":
- return list(self.traversal_search(query, **kwargs))
- elif search_type == "mmr_traversal":
- return list(self.mmr_traversal_search(query, **kwargs))
- else:
- raise ValueError(
- f"search_type of {search_type} not allowed. Expected "
- "search_type to be 'similarity', 'similarity_score_threshold', "
- "'mmr', 'traversal', or 'mmr_traversal'."
- )
-
- async def asearch(
- self, query: str, search_type: str, **kwargs: Any
- ) -> list[Document]:
- if search_type == "similarity":
- return await self.asimilarity_search(query, **kwargs)
- elif search_type == "similarity_score_threshold":
- docs_and_similarities = await self.asimilarity_search_with_relevance_scores(
- query, **kwargs
- )
- return [doc for doc, _ in docs_and_similarities]
- elif search_type == "mmr":
- return await self.amax_marginal_relevance_search(query, **kwargs)
- elif search_type == "traversal":
- return [doc async for doc in self.atraversal_search(query, **kwargs)]
- elif search_type == "mmr_traversal":
- return [doc async for doc in self.ammr_traversal_search(query, **kwargs)]
- else:
- raise ValueError(
- f"search_type of {search_type} not allowed. Expected "
- "search_type to be 'similarity', 'similarity_score_threshold', "
- "'mmr', 'traversal', or 'mmr_traversal'."
- )
-
- def as_retriever(self, **kwargs: Any) -> GraphVectorStoreRetriever:
- """Return GraphVectorStoreRetriever initialized from this GraphVectorStore.
-
- Args:
- **kwargs: Keyword arguments to pass to the search function.
- Can include:
-
- - search_type (Optional[str]): Defines the type of search that
- the Retriever should perform.
- Can be ``traversal`` (default), ``similarity``, ``mmr``,
- ``mmr_traversal``, or ``similarity_score_threshold``.
- - search_kwargs (Optional[Dict]): Keyword arguments to pass to the
- search function. Can include things like:
-
- - k(int): Amount of documents to return (Default: 4).
- - depth(int): The maximum depth of edges to traverse (Default: 1).
- Only applies to search_type: ``traversal`` and ``mmr_traversal``.
- - score_threshold(float): Minimum relevance threshold
- for similarity_score_threshold.
- - fetch_k(int): Amount of documents to pass to MMR algorithm
- (Default: 20).
- - lambda_mult(float): Diversity of results returned by MMR;
- 1 for minimum diversity and 0 for maximum. (Default: 0.5).
- Returns:
- Retriever for this GraphVectorStore.
-
- Examples:
-
- .. code-block:: python
-
- # Retrieve documents traversing edges
- docsearch.as_retriever(
- search_type="traversal",
- search_kwargs={'k': 6, 'depth': 2}
- )
-
- # Retrieve documents with higher diversity
- # Useful if your dataset has many similar documents
- docsearch.as_retriever(
- search_type="mmr_traversal",
- search_kwargs={'k': 6, 'lambda_mult': 0.25, 'depth': 2}
- )
-
- # Fetch more documents for the MMR algorithm to consider
- # But only return the top 5
- docsearch.as_retriever(
- search_type="mmr_traversal",
- search_kwargs={'k': 5, 'fetch_k': 50, 'depth': 2}
- )
-
- # Only retrieve documents that have a relevance score
- # Above a certain threshold
- docsearch.as_retriever(
- search_type="similarity_score_threshold",
- search_kwargs={'score_threshold': 0.8}
- )
-
- # Only get the single most similar document from the dataset
- docsearch.as_retriever(search_kwargs={'k': 1})
-
- """
- return GraphVectorStoreRetriever(vectorstore=self, **kwargs)
-
-
-@deprecated(
- since="0.3.21",
- removal="0.5",
- addendum=DEPRECATION_ADDENDUM,
-)
-class GraphVectorStoreRetriever(VectorStoreRetriever):
- """Retriever for GraphVectorStore.
-
- A graph vector store retriever is a retriever that uses a graph vector store to
- retrieve documents.
- It is similar to a vector store retriever, except that it uses both vector
- similarity and graph connections to retrieve documents.
- It uses the search methods implemented by a graph vector store, like traversal
- search and MMR traversal search, to query the texts in the graph vector store.
-
- Example::
-
- store = CassandraGraphVectorStore(...)
- retriever = store.as_retriever()
- retriever.invoke("What is ...")
-
- .. seealso::
-
- :mod:`How to use a graph vector store `
-
- How to use a graph vector store as a retriever
- ==============================================
-
- Creating a retriever from a graph vector store
- ----------------------------------------------
-
- You can build a retriever from a graph vector store using its
- :meth:`~langchain_community.graph_vectorstores.base.GraphVectorStore.as_retriever`
- method.
-
- First we instantiate a graph vector store.
- We will use a store backed by Cassandra
- :class:`~langchain_community.graph_vectorstores.cassandra.CassandraGraphVectorStore`
- graph vector store::
-
- from langchain_community.document_loaders import TextLoader
- from langchain_community.graph_vectorstores import CassandraGraphVectorStore
- from langchain_community.graph_vectorstores.extractors import (
- KeybertLinkExtractor,
- LinkExtractorTransformer,
- )
- from langchain_openai import OpenAIEmbeddings
- from langchain_text_splitters import CharacterTextSplitter
-
- loader = TextLoader("state_of_the_union.txt")
- documents = loader.load()
-
- text_splitter = CharacterTextSplitter(chunk_size=1000, chunk_overlap=0)
- texts = text_splitter.split_documents(documents)
-
- pipeline = LinkExtractorTransformer([KeybertLinkExtractor()])
- pipeline.transform_documents(texts)
- embeddings = OpenAIEmbeddings()
- graph_vectorstore = CassandraGraphVectorStore.from_documents(texts, embeddings)
-
- We can then instantiate a retriever::
-
- retriever = graph_vectorstore.as_retriever()
-
- This creates a retriever (specifically a ``GraphVectorStoreRetriever``), which we
- can use in the usual way::
-
- docs = retriever.invoke("what did the president say about ketanji brown jackson?")
-
- Maximum marginal relevance traversal retrieval
- ----------------------------------------------
-
- By default, the graph vector store retriever uses similarity search, then expands
- the retrieved set by following a fixed number of graph edges.
- If the underlying graph vector store supports maximum marginal relevance traversal,
- you can specify that as the search type.
-
- MMR-traversal is a retrieval method combining MMR and graph traversal.
- The strategy first retrieves the top fetch_k results by similarity to the question.
- It then iteratively expands the set of fetched documents by following adjacent_k
- graph edges and selects the top k results based on maximum-marginal relevance using
- the given ``lambda_mult``::
-
- retriever = graph_vectorstore.as_retriever(search_type="mmr_traversal")
-
- Passing search parameters
- -------------------------
-
- We can pass parameters to the underlying graph vector store's search methods using
- ``search_kwargs``.
-
- Specifying graph traversal depth
- ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
-
- For example, we can set the graph traversal depth to only return documents
- reachable through a given number of graph edges::
-
- retriever = graph_vectorstore.as_retriever(search_kwargs={"depth": 3})
-
- Specifying MMR parameters
- ^^^^^^^^^^^^^^^^^^^^^^^^^
-
- When using search type ``mmr_traversal``, several parameters of the MMR algorithm
- can be configured.
-
- The ``fetch_k`` parameter determines how many documents are fetched using vector
- similarity and ``adjacent_k`` parameter determines how many documents are fetched
- using graph edges.
- The ``lambda_mult`` parameter controls how the MMR re-ranking weights similarity to
- the query string vs diversity among the retrieved documents as fetched documents
- are selected for the set of ``k`` final results::
-
- retriever = graph_vectorstore.as_retriever(
- search_type="mmr",
- search_kwargs={"fetch_k": 20, "adjacent_k": 20, "lambda_mult": 0.25},
- )
-
- Specifying top k
- ^^^^^^^^^^^^^^^^
-
- We can also limit the number of documents ``k`` returned by the retriever.
-
- Note that if ``depth`` is greater than zero, the retriever may return more documents
- than is specified by ``k``, since both the original ``k`` documents retrieved using
- vector similarity and any documents connected via graph edges will be returned::
-
- retriever = graph_vectorstore.as_retriever(search_kwargs={"k": 1})
-
- Similarity score threshold retrieval
- ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
-
- For example, we can set a similarity score threshold and only return documents with
- a score above that threshold::
-
- retriever = graph_vectorstore.as_retriever(search_kwargs={"score_threshold": 0.5})
- """ # noqa: E501
-
- vectorstore: VectorStore
- """VectorStore to use for retrieval."""
- search_type: str = "traversal"
- """Type of search to perform. Defaults to "traversal"."""
- allowed_search_types: ClassVar[Collection[str]] = (
- "similarity",
- "similarity_score_threshold",
- "mmr",
- "traversal",
- "mmr_traversal",
- )
-
- @property
- def graph_vectorstore(self) -> GraphVectorStore:
- return cast(GraphVectorStore, self.vectorstore)
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun, **kwargs: Any
- ) -> list[Document]:
- if self.search_type == "traversal":
- return list(
- self.graph_vectorstore.traversal_search(query, **self.search_kwargs)
- )
- elif self.search_type == "mmr_traversal":
- return list(
- self.graph_vectorstore.mmr_traversal_search(query, **self.search_kwargs)
- )
- else:
- return super()._get_relevant_documents(query, run_manager=run_manager)
-
- async def _aget_relevant_documents(
- self,
- query: str,
- *,
- run_manager: AsyncCallbackManagerForRetrieverRun,
- **kwargs: Any,
- ) -> list[Document]:
- if self.search_type == "traversal":
- return [
- doc
- async for doc in self.graph_vectorstore.atraversal_search(
- query, **self.search_kwargs
- )
- ]
- elif self.search_type == "mmr_traversal":
- return [
- doc
- async for doc in self.graph_vectorstore.ammr_traversal_search(
- query, **self.search_kwargs
- )
- ]
- else:
- return await super()._aget_relevant_documents(
- query, run_manager=run_manager
- )
diff --git a/libs/community/langchain_community/graph_vectorstores/cassandra.py b/libs/community/langchain_community/graph_vectorstores/cassandra.py
deleted file mode 100644
index 5b377d61f5..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/cassandra.py
+++ /dev/null
@@ -1,1268 +0,0 @@
-"""Apache Cassandra DB graph vector store integration."""
-
-from __future__ import annotations
-
-import asyncio
-import json
-import logging
-import secrets
-from dataclasses import asdict, is_dataclass
-from typing import (
- TYPE_CHECKING,
- Any,
- AsyncIterable,
- Iterable,
- List,
- Optional,
- Sequence,
- Tuple,
- Type,
- TypeVar,
- cast,
-)
-
-from langchain_core._api import beta
-from langchain_core.documents import Document
-from typing_extensions import override
-
-from langchain_community.graph_vectorstores.base import GraphVectorStore, Node
-from langchain_community.graph_vectorstores.links import METADATA_LINKS_KEY, Link
-from langchain_community.graph_vectorstores.mmr_helper import MmrHelper
-from langchain_community.utilities.cassandra import SetupMode
-from langchain_community.vectorstores.cassandra import Cassandra as CassandraVectorStore
-
-CGVST = TypeVar("CGVST", bound="CassandraGraphVectorStore")
-
-if TYPE_CHECKING:
- from cassandra.cluster import Session
- from langchain_core.embeddings import Embeddings
-
-
-logger = logging.getLogger(__name__)
-
-
-class AdjacentNode:
- id: str
- links: list[Link]
- embedding: list[float]
-
- def __init__(self, node: Node, embedding: list[float]) -> None:
- """Create an Adjacent Node."""
- self.id = node.id or ""
- self.links = node.links
- self.embedding = embedding
-
-
-def _serialize_links(links: list[Link]) -> str:
- class SetAndLinkEncoder(json.JSONEncoder):
- def default(self, obj: Any) -> Any: # noqa: ANN401
- if not isinstance(obj, type) and is_dataclass(obj):
- return asdict(obj)
-
- if isinstance(obj, Iterable):
- return list(obj)
-
- # Let the base class default method raise the TypeError
- return super().default(obj)
-
- return json.dumps(links, cls=SetAndLinkEncoder)
-
-
-def _deserialize_links(json_blob: str | None) -> set[Link]:
- return {
- Link(kind=link["kind"], direction=link["direction"], tag=link["tag"])
- for link in cast(list[dict[str, Any]], json.loads(json_blob or "[]"))
- }
-
-
-def _metadata_link_key(link: Link) -> str:
- return f"link:{link.kind}:{link.tag}"
-
-
-def _metadata_link_value() -> str:
- return "link"
-
-
-def _doc_to_node(doc: Document) -> Node:
- metadata = doc.metadata.copy()
- links = _deserialize_links(metadata.get(METADATA_LINKS_KEY))
- metadata[METADATA_LINKS_KEY] = links
-
- return Node(
- id=doc.id,
- text=doc.page_content,
- metadata=metadata,
- links=list(links),
- )
-
-
-def _incoming_links(node: Node | AdjacentNode) -> set[Link]:
- return {link for link in node.links if link.direction in ["in", "bidir"]}
-
-
-def _outgoing_links(node: Node | AdjacentNode) -> set[Link]:
- return {link for link in node.links if link.direction in ["out", "bidir"]}
-
-
-@beta()
-class CassandraGraphVectorStore(GraphVectorStore):
- def __init__(
- self,
- embedding: Embeddings,
- session: Session | None = None,
- keyspace: str | None = None,
- table_name: str = "",
- ttl_seconds: int | None = None,
- *,
- body_index_options: list[tuple[str, Any]] | None = None,
- setup_mode: SetupMode = SetupMode.SYNC,
- metadata_deny_list: Optional[list[str]] = None,
- ) -> None:
- """Apache Cassandra(R) for graph-vector-store workloads.
-
- To use it, you need a recent installation of the `cassio` library
- and a Cassandra cluster / Astra DB instance supporting vector capabilities.
-
- Example:
- .. code-block:: python
-
- from langchain_community.graph_vectorstores import
- CassandraGraphVectorStore
- from langchain_openai import OpenAIEmbeddings
-
- embeddings = OpenAIEmbeddings()
- session = ... # create your Cassandra session object
- keyspace = 'my_keyspace' # the keyspace should exist already
- table_name = 'my_graph_vector_store'
- vectorstore = CassandraGraphVectorStore(
- embeddings,
- session,
- keyspace,
- table_name,
- )
-
- Args:
- embedding: Embedding function to use.
- session: Cassandra driver session. If not provided, it is resolved from
- cassio.
- keyspace: Cassandra keyspace. If not provided, it is resolved from cassio.
- table_name: Cassandra table (required).
- ttl_seconds: Optional time-to-live for the added texts.
- body_index_options: Optional options used to create the body index.
- Eg. body_index_options = [cassio.table.cql.STANDARD_ANALYZER]
- setup_mode: mode used to create the Cassandra table (SYNC,
- ASYNC or OFF).
- metadata_deny_list: Optional list of metadata keys to not index.
- i.e. to fine-tune which of the metadata fields are indexed.
- Note: if you plan to have massive unique text metadata entries,
- consider not indexing them for performance
- (and to overcome max-length limitations).
- Note: the `metadata_indexing` parameter from
- langchain_community.utilities.cassandra.Cassandra is not
- exposed since CassandraGraphVectorStore only supports the
- deny_list option.
- """
- self.embedding = embedding
-
- if metadata_deny_list is None:
- metadata_deny_list = []
- metadata_deny_list.append(METADATA_LINKS_KEY)
-
- self.vector_store = CassandraVectorStore(
- embedding=embedding,
- session=session,
- keyspace=keyspace,
- table_name=table_name,
- ttl_seconds=ttl_seconds,
- body_index_options=body_index_options,
- setup_mode=setup_mode,
- metadata_indexing=("deny_list", metadata_deny_list),
- )
-
- store_session: Session = self.vector_store.session
-
- self._insert_node = store_session.prepare(
- f"""
- INSERT INTO {keyspace}.{table_name} (
- row_id, body_blob, vector, attributes_blob, metadata_s
- ) VALUES (?, ?, ?, ?, ?)
- """ # noqa: S608
- )
-
- @property
- @override
- def embeddings(self) -> Embeddings | None:
- return self.embedding
-
- def _get_metadata_filter(
- self,
- metadata: dict[str, Any] | None = None,
- outgoing_link: Link | None = None,
- ) -> dict[str, Any]:
- if outgoing_link is None:
- return metadata or {}
-
- metadata_filter = {} if metadata is None else metadata.copy()
- metadata_filter[_metadata_link_key(link=outgoing_link)] = _metadata_link_value()
- return metadata_filter
-
- def _restore_links(self, doc: Document) -> Document:
- """Restores the links in the document by deserializing them from metadata.
-
- Args:
- doc: A single Document
-
- Returns:
- The same Document with restored links.
- """
- links = _deserialize_links(doc.metadata.get(METADATA_LINKS_KEY))
- doc.metadata[METADATA_LINKS_KEY] = links
- # TODO: Could this be skipped if we put these metadata entries
- # only in the searchable `metadata_s` column?
- for incoming_link_key in [
- _metadata_link_key(link=link)
- for link in links
- if link.direction in ["in", "bidir"]
- ]:
- if incoming_link_key in doc.metadata:
- del doc.metadata[incoming_link_key]
-
- return doc
-
- def _get_node_metadata_for_insertion(self, node: Node) -> dict[str, Any]:
- metadata = node.metadata.copy()
- metadata[METADATA_LINKS_KEY] = _serialize_links(node.links)
- # TODO: Could we could put these metadata entries
- # only in the searchable `metadata_s` column?
- for incoming_link in _incoming_links(node=node):
- metadata[_metadata_link_key(link=incoming_link)] = _metadata_link_value()
- return metadata
-
- def _get_docs_for_insertion(
- self, nodes: Iterable[Node]
- ) -> tuple[list[Document], list[str]]:
- docs = []
- ids = []
- for node in nodes:
- node_id = secrets.token_hex(8) if not node.id else node.id
-
- doc = Document(
- page_content=node.text,
- metadata=self._get_node_metadata_for_insertion(node=node),
- id=node_id,
- )
- docs.append(doc)
- ids.append(node_id)
- return (docs, ids)
-
- @override
- def add_nodes(
- self,
- nodes: Iterable[Node],
- **kwargs: Any,
- ) -> Iterable[str]:
- """Add nodes to the graph store.
-
- Args:
- nodes: the nodes to add.
- **kwargs: Additional keyword arguments.
- """
- (docs, ids) = self._get_docs_for_insertion(nodes=nodes)
- return self.vector_store.add_documents(docs, ids=ids)
-
- @override
- async def aadd_nodes(
- self,
- nodes: Iterable[Node],
- **kwargs: Any,
- ) -> AsyncIterable[str]:
- """Add nodes to the graph store.
-
- Args:
- nodes: the nodes to add.
- **kwargs: Additional keyword arguments.
- """
- (docs, ids) = self._get_docs_for_insertion(nodes=nodes)
- for inserted_id in await self.vector_store.aadd_documents(docs, ids=ids):
- yield inserted_id
-
- @override
- def similarity_search(
- self,
- query: str,
- k: int = 4,
- filter: dict[str, Any] | None = None,
- **kwargs: Any,
- ) -> list[Document]:
- """Retrieve documents from this graph store.
-
- Args:
- query: The query string.
- k: The number of Documents to return. Defaults to 4.
- filter: Optional metadata to filter the results.
- **kwargs: Additional keyword arguments.
-
- Returns:
- Collection of retrieved documents.
- """
- return [
- self._restore_links(doc)
- for doc in self.vector_store.similarity_search(
- query=query,
- k=k,
- filter=filter,
- **kwargs,
- )
- ]
-
- @override
- async def asimilarity_search(
- self,
- query: str,
- k: int = 4,
- filter: dict[str, Any] | None = None,
- **kwargs: Any,
- ) -> list[Document]:
- """Retrieve documents from this graph store.
-
- Args:
- query: The query string.
- k: The number of Documents to return. Defaults to 4.
- filter: Optional metadata to filter the results.
- **kwargs: Additional keyword arguments.
-
- Returns:
- Collection of retrieved documents.
- """
- return [
- self._restore_links(doc)
- for doc in await self.vector_store.asimilarity_search(
- query=query,
- k=k,
- filter=filter,
- **kwargs,
- )
- ]
-
- @override
- def similarity_search_by_vector(
- self,
- embedding: list[float],
- k: int = 4,
- filter: dict[str, Any] | None = None,
- **kwargs: Any,
- ) -> list[Document]:
- """Return docs most similar to embedding vector.
-
- Args:
- embedding: Embedding to look up documents similar to.
- k: Number of Documents to return. Defaults to 4.
- filter: Filter on the metadata to apply.
- **kwargs: Additional arguments are ignored.
-
- Returns:
- The list of Documents most similar to the query vector.
- """
- return [
- self._restore_links(doc)
- for doc in self.vector_store.similarity_search_by_vector(
- embedding,
- k=k,
- filter=filter,
- **kwargs,
- )
- ]
-
- @override
- async def asimilarity_search_by_vector(
- self,
- embedding: list[float],
- k: int = 4,
- filter: dict[str, Any] | None = None,
- **kwargs: Any,
- ) -> list[Document]:
- """Return docs most similar to embedding vector.
-
- Args:
- embedding: Embedding to look up documents similar to.
- k: Number of Documents to return. Defaults to 4.
- filter: Filter on the metadata to apply.
- **kwargs: Additional arguments are ignored.
-
- Returns:
- The list of Documents most similar to the query vector.
- """
- return [
- self._restore_links(doc)
- for doc in await self.vector_store.asimilarity_search_by_vector(
- embedding,
- k=k,
- filter=filter,
- **kwargs,
- )
- ]
-
- def metadata_search(
- self,
- filter: dict[str, Any] | None = None, # noqa: A002
- n: int = 5,
- ) -> Iterable[Document]:
- """Get documents via a metadata search.
-
- Args:
- filter: the metadata to query for.
- n: the maximum number of documents to return.
- """
- return [
- self._restore_links(doc)
- for doc in self.vector_store.metadata_search(
- filter=filter or {},
- n=n,
- )
- ]
-
- async def ametadata_search(
- self,
- filter: dict[str, Any] | None = None, # noqa: A002
- n: int = 5,
- ) -> Iterable[Document]:
- """Get documents via a metadata search.
-
- Args:
- filter: the metadata to query for.
- n: the maximum number of documents to return.
- """
- return [
- self._restore_links(doc)
- for doc in await self.vector_store.ametadata_search(
- filter=filter or {},
- n=n,
- )
- ]
-
- def get_by_document_id(self, document_id: str) -> Document | None:
- """Retrieve a single document from the store, given its document ID.
-
- Args:
- document_id: The document ID
-
- Returns:
- The the document if it exists. Otherwise None.
- """
- doc = self.vector_store.get_by_document_id(document_id=document_id)
- return self._restore_links(doc) if doc is not None else None
-
- async def aget_by_document_id(self, document_id: str) -> Document | None:
- """Retrieve a single document from the store, given its document ID.
-
- Args:
- document_id: The document ID
-
- Returns:
- The the document if it exists. Otherwise None.
- """
- doc = await self.vector_store.aget_by_document_id(document_id=document_id)
- return self._restore_links(doc) if doc is not None else None
-
- def get_node(self, node_id: str) -> Node | None:
- """Retrieve a single node from the store, given its ID.
-
- Args:
- node_id: The node ID
-
- Returns:
- The the node if it exists. Otherwise None.
- """
- doc = self.vector_store.get_by_document_id(document_id=node_id)
- if doc is None:
- return None
- return _doc_to_node(doc=doc)
-
- @override
- async def ammr_traversal_search( # noqa: C901
- self,
- query: str,
- *,
- initial_roots: Sequence[str] = (),
- k: int = 4,
- depth: int = 2,
- fetch_k: int = 100,
- adjacent_k: int = 10,
- lambda_mult: float = 0.5,
- score_threshold: float = float("-inf"),
- filter: dict[str, Any] | None = None,
- **kwargs: Any,
- ) -> AsyncIterable[Document]:
- """Retrieve documents from this graph store using MMR-traversal.
-
- This strategy first retrieves the top `fetch_k` results by similarity to
- the question. It then selects the top `k` results based on
- maximum-marginal relevance using the given `lambda_mult`.
-
- At each step, it considers the (remaining) documents from `fetch_k` as
- well as any documents connected by edges to a selected document
- retrieved based on similarity (a "root").
-
- Args:
- query: The query string to search for.
- initial_roots: Optional list of document IDs to use for initializing search.
- The top `adjacent_k` nodes adjacent to each initial root will be
- included in the set of initial candidates. To fetch only in the
- neighborhood of these nodes, set `fetch_k = 0`.
- k: Number of Documents to return. Defaults to 4.
- fetch_k: Number of initial Documents to fetch via similarity.
- Will be added to the nodes adjacent to `initial_roots`.
- Defaults to 100.
- adjacent_k: Number of adjacent Documents to fetch.
- Defaults to 10.
- depth: Maximum depth of a node (number of edges) from a node
- retrieved via similarity. Defaults to 2.
- lambda_mult: Number between 0 and 1 that determines the degree
- of diversity among the results with 0 corresponding to maximum
- diversity and 1 to minimum diversity. Defaults to 0.5.
- score_threshold: Only documents with a score greater than or equal
- this threshold will be chosen. Defaults to -infinity.
- filter: Optional metadata to filter the results.
- **kwargs: Additional keyword arguments.
- """
- query_embedding = self.embedding.embed_query(query)
- helper = MmrHelper(
- k=k,
- query_embedding=query_embedding,
- lambda_mult=lambda_mult,
- score_threshold=score_threshold,
- )
-
- # For each unselected node, stores the outgoing links.
- outgoing_links_map: dict[str, set[Link]] = {}
- visited_links: set[Link] = set()
- # Map from id to Document
- retrieved_docs: dict[str, Document] = {}
-
- async def fetch_neighborhood(neighborhood: Sequence[str]) -> None:
- nonlocal outgoing_links_map, visited_links, retrieved_docs
-
- # Put the neighborhood into the outgoing links, to avoid adding it
- # to the candidate set in the future.
- outgoing_links_map.update(
- {content_id: set() for content_id in neighborhood}
- )
-
- # Initialize the visited_links with the set of outgoing links from the
- # neighborhood. This prevents re-visiting them.
- visited_links = await self._get_outgoing_links(neighborhood)
-
- # Call `self._get_adjacent` to fetch the candidates.
- adjacent_nodes = await self._get_adjacent(
- links=visited_links,
- query_embedding=query_embedding,
- k_per_link=adjacent_k,
- filter=filter,
- retrieved_docs=retrieved_docs,
- )
-
- new_candidates: dict[str, list[float]] = {}
- for adjacent_node in adjacent_nodes:
- if adjacent_node.id not in outgoing_links_map:
- outgoing_links_map[adjacent_node.id] = _outgoing_links(
- node=adjacent_node
- )
- new_candidates[adjacent_node.id] = adjacent_node.embedding
- helper.add_candidates(new_candidates)
-
- async def fetch_initial_candidates() -> None:
- nonlocal outgoing_links_map, visited_links, retrieved_docs
-
- results = (
- await self.vector_store.asimilarity_search_with_embedding_id_by_vector(
- embedding=query_embedding,
- k=fetch_k,
- filter=filter,
- )
- )
-
- candidates: dict[str, list[float]] = {}
- for doc, embedding, doc_id in results:
- if doc_id not in retrieved_docs:
- retrieved_docs[doc_id] = doc
-
- if doc_id not in outgoing_links_map:
- node = _doc_to_node(doc)
- outgoing_links_map[doc_id] = _outgoing_links(node=node)
- candidates[doc_id] = embedding
- helper.add_candidates(candidates)
-
- if initial_roots:
- await fetch_neighborhood(initial_roots)
- if fetch_k > 0:
- await fetch_initial_candidates()
-
- # Tracks the depth of each candidate.
- depths = {candidate_id: 0 for candidate_id in helper.candidate_ids()}
-
- # Select the best item, K times.
- for _ in range(k):
- selected_id = helper.pop_best()
-
- if selected_id is None:
- break
-
- next_depth = depths[selected_id] + 1
- if next_depth < depth:
- # If the next nodes would not exceed the depth limit, find the
- # adjacent nodes.
-
- # Find the links linked to from the selected ID.
- selected_outgoing_links = outgoing_links_map.pop(selected_id)
-
- # Don't re-visit already visited links.
- selected_outgoing_links.difference_update(visited_links)
-
- # Find the nodes with incoming links from those links.
- adjacent_nodes = await self._get_adjacent(
- links=selected_outgoing_links,
- query_embedding=query_embedding,
- k_per_link=adjacent_k,
- filter=filter,
- retrieved_docs=retrieved_docs,
- )
-
- # Record the selected_outgoing_links as visited.
- visited_links.update(selected_outgoing_links)
-
- new_candidates = {}
- for adjacent_node in adjacent_nodes:
- if adjacent_node.id not in outgoing_links_map:
- outgoing_links_map[adjacent_node.id] = _outgoing_links(
- node=adjacent_node
- )
- new_candidates[adjacent_node.id] = adjacent_node.embedding
- if next_depth < depths.get(adjacent_node.id, depth + 1):
- # If this is a new shortest depth, or there was no
- # previous depth, update the depths. This ensures that
- # when we discover a node we will have the shortest
- # depth available.
- #
- # NOTE: No effort is made to traverse from nodes that
- # were previously selected if they become reachable via
- # a shorter path via nodes selected later. This is
- # currently "intended", but may be worth experimenting
- # with.
- depths[adjacent_node.id] = next_depth
- helper.add_candidates(new_candidates)
-
- for doc_id, similarity_score, mmr_score in zip(
- helper.selected_ids,
- helper.selected_similarity_scores,
- helper.selected_mmr_scores,
- ):
- if doc_id in retrieved_docs:
- doc = self._restore_links(retrieved_docs[doc_id])
- doc.metadata["similarity_score"] = similarity_score
- doc.metadata["mmr_score"] = mmr_score
- yield doc
- else:
- msg = f"retrieved_docs should contain id: {doc_id}"
- raise RuntimeError(msg)
-
- @override
- def mmr_traversal_search(
- self,
- query: str,
- *,
- initial_roots: Sequence[str] = (),
- k: int = 4,
- depth: int = 2,
- fetch_k: int = 100,
- adjacent_k: int = 10,
- lambda_mult: float = 0.5,
- score_threshold: float = float("-inf"),
- filter: dict[str, Any] | None = None,
- **kwargs: Any,
- ) -> Iterable[Document]:
- """Retrieve documents from this graph store using MMR-traversal.
-
- This strategy first retrieves the top `fetch_k` results by similarity to
- the question. It then selects the top `k` results based on
- maximum-marginal relevance using the given `lambda_mult`.
-
- At each step, it considers the (remaining) documents from `fetch_k` as
- well as any documents connected by edges to a selected document
- retrieved based on similarity (a "root").
-
- Args:
- query: The query string to search for.
- initial_roots: Optional list of document IDs to use for initializing search.
- The top `adjacent_k` nodes adjacent to each initial root will be
- included in the set of initial candidates. To fetch only in the
- neighborhood of these nodes, set `fetch_k = 0`.
- k: Number of Documents to return. Defaults to 4.
- fetch_k: Number of initial Documents to fetch via similarity.
- Will be added to the nodes adjacent to `initial_roots`.
- Defaults to 100.
- adjacent_k: Number of adjacent Documents to fetch.
- Defaults to 10.
- depth: Maximum depth of a node (number of edges) from a node
- retrieved via similarity. Defaults to 2.
- lambda_mult: Number between 0 and 1 that determines the degree
- of diversity among the results with 0 corresponding to maximum
- diversity and 1 to minimum diversity. Defaults to 0.5.
- score_threshold: Only documents with a score greater than or equal
- this threshold will be chosen. Defaults to -infinity.
- filter: Optional metadata to filter the results.
- **kwargs: Additional keyword arguments.
- """
-
- async def collect_docs() -> Iterable[Document]:
- async_iter = self.ammr_traversal_search(
- query=query,
- initial_roots=initial_roots,
- k=k,
- depth=depth,
- fetch_k=fetch_k,
- adjacent_k=adjacent_k,
- lambda_mult=lambda_mult,
- score_threshold=score_threshold,
- filter=filter,
- **kwargs,
- )
- return [doc async for doc in async_iter]
-
- return asyncio.run(collect_docs())
-
- @override
- async def atraversal_search( # noqa: C901
- self,
- query: str,
- *,
- k: int = 4,
- depth: int = 1,
- filter: dict[str, Any] | None = None,
- **kwargs: Any,
- ) -> AsyncIterable[Document]:
- """Retrieve documents from this knowledge store.
-
- First, `k` nodes are retrieved using a vector search for the `query` string.
- Then, additional nodes are discovered up to the given `depth` from those
- starting nodes.
-
- Args:
- query: The query string.
- k: The number of Documents to return from the initial vector search.
- Defaults to 4.
- depth: The maximum depth of edges to traverse. Defaults to 1.
- filter: Optional metadata to filter the results.
- **kwargs: Additional keyword arguments.
-
- Returns:
- Collection of retrieved documents.
- """
- # Depth 0:
- # Query for `k` nodes similar to the question.
- # Retrieve `content_id` and `outgoing_links()`.
- #
- # Depth 1:
- # Query for nodes that have an incoming link in the `outgoing_links()` set.
- # Combine node IDs.
- # Query for `outgoing_links()` of those "new" node IDs.
- #
- # ...
-
- # Map from visited ID to depth
- visited_ids: dict[str, int] = {}
-
- # Map from visited link to depth
- visited_links: dict[Link, int] = {}
-
- # Map from id to Document
- retrieved_docs: dict[str, Document] = {}
-
- async def visit_nodes(d: int, docs: Iterable[Document]) -> None:
- """Recursively visit nodes and their outgoing links."""
- nonlocal visited_ids, visited_links, retrieved_docs
-
- # Iterate over nodes, tracking the *new* outgoing links for this
- # depth. These are links that are either new, or newly discovered at a
- # lower depth.
- outgoing_links: set[Link] = set()
- for doc in docs:
- if doc.id is not None:
- if doc.id not in retrieved_docs:
- retrieved_docs[doc.id] = doc
-
- # If this node is at a closer depth, update visited_ids
- if d <= visited_ids.get(doc.id, depth):
- visited_ids[doc.id] = d
-
- # If we can continue traversing from this node,
- if d < depth:
- node = _doc_to_node(doc=doc)
- # Record any new (or newly discovered at a lower depth)
- # links to the set to traverse.
- for link in _outgoing_links(node=node):
- if d <= visited_links.get(link, depth):
- # Record that we'll query this link at the
- # given depth, so we don't fetch it again
- # (unless we find it an earlier depth)
- visited_links[link] = d
- outgoing_links.add(link)
-
- if outgoing_links:
- metadata_search_tasks = []
- for outgoing_link in outgoing_links:
- metadata_filter = self._get_metadata_filter(
- metadata=filter,
- outgoing_link=outgoing_link,
- )
- metadata_search_tasks.append(
- asyncio.create_task(
- self.vector_store.ametadata_search(
- filter=metadata_filter, n=1000
- )
- )
- )
- results = await asyncio.gather(*metadata_search_tasks)
-
- # Visit targets concurrently
- visit_target_tasks = [
- visit_targets(d=d + 1, docs=docs) for docs in results
- ]
- await asyncio.gather(*visit_target_tasks)
-
- async def visit_targets(d: int, docs: Iterable[Document]) -> None:
- """Visit target nodes retrieved from outgoing links."""
- nonlocal visited_ids, retrieved_docs
-
- new_ids_at_next_depth = set()
- for doc in docs:
- if doc.id is not None:
- if doc.id not in retrieved_docs:
- retrieved_docs[doc.id] = doc
-
- if d <= visited_ids.get(doc.id, depth):
- new_ids_at_next_depth.add(doc.id)
-
- if new_ids_at_next_depth:
- visit_node_tasks = [
- visit_nodes(d=d, docs=[retrieved_docs[doc_id]])
- for doc_id in new_ids_at_next_depth
- if doc_id in retrieved_docs
- ]
-
- fetch_tasks = [
- asyncio.create_task(
- self.vector_store.aget_by_document_id(document_id=doc_id)
- )
- for doc_id in new_ids_at_next_depth
- if doc_id not in retrieved_docs
- ]
-
- new_docs: list[Document | None] = await asyncio.gather(*fetch_tasks)
-
- visit_node_tasks.extend(
- visit_nodes(d=d, docs=[new_doc])
- for new_doc in new_docs
- if new_doc is not None
- )
-
- await asyncio.gather(*visit_node_tasks)
-
- # Start the traversal
- initial_docs = self.vector_store.similarity_search(
- query=query,
- k=k,
- filter=filter,
- )
- await visit_nodes(d=0, docs=initial_docs)
-
- for doc_id in visited_ids:
- if doc_id in retrieved_docs:
- yield self._restore_links(retrieved_docs[doc_id])
- else:
- msg = f"retrieved_docs should contain id: {doc_id}"
- raise RuntimeError(msg)
-
- @override
- def traversal_search(
- self,
- query: str,
- *,
- k: int = 4,
- depth: int = 1,
- filter: dict[str, Any] | None = None,
- **kwargs: Any,
- ) -> Iterable[Document]:
- """Retrieve documents from this knowledge store.
-
- First, `k` nodes are retrieved using a vector search for the `query` string.
- Then, additional nodes are discovered up to the given `depth` from those
- starting nodes.
-
- Args:
- query: The query string.
- k: The number of Documents to return from the initial vector search.
- Defaults to 4.
- depth: The maximum depth of edges to traverse. Defaults to 1.
- filter: Optional metadata to filter the results.
- **kwargs: Additional keyword arguments.
-
- Returns:
- Collection of retrieved documents.
- """
-
- async def collect_docs() -> Iterable[Document]:
- async_iter = self.atraversal_search(
- query=query,
- k=k,
- depth=depth,
- filter=filter,
- **kwargs,
- )
- return [doc async for doc in async_iter]
-
- return asyncio.run(collect_docs())
-
- async def _get_outgoing_links(self, source_ids: Iterable[str]) -> set[Link]:
- """Return the set of outgoing links for the given source IDs asynchronously.
-
- Args:
- source_ids: The IDs of the source nodes to retrieve outgoing links for.
-
- Returns:
- A set of `Link` objects representing the outgoing links from the source
- nodes.
- """
- links = set()
-
- # Create coroutine objects without scheduling them yet
- coroutines = [
- self.vector_store.aget_by_document_id(document_id=source_id)
- for source_id in source_ids
- ]
-
- # Schedule and await all coroutines
- docs = await asyncio.gather(*coroutines)
-
- for doc in docs:
- if doc is not None:
- node = _doc_to_node(doc=doc)
- links.update(_outgoing_links(node=node))
-
- return links
-
- async def _get_adjacent(
- self,
- links: set[Link],
- query_embedding: list[float],
- retrieved_docs: dict[str, Document],
- k_per_link: int | None = None,
- filter: dict[str, Any] | None = None, # noqa: A002
- ) -> Iterable[AdjacentNode]:
- """Return the target nodes with incoming links from any of the given links.
-
- Args:
- links: The links to look for.
- query_embedding: The query embedding. Used to rank target nodes.
- retrieved_docs: A cache of retrieved docs. This will be added to.
- k_per_link: The number of target nodes to fetch for each link.
- filter: Optional metadata to filter the results.
-
- Returns:
- Iterable of adjacent edges.
- """
- targets: dict[str, AdjacentNode] = {}
-
- tasks = []
- for link in links:
- metadata_filter = self._get_metadata_filter(
- metadata=filter,
- outgoing_link=link,
- )
-
- tasks.append(
- self.vector_store.asimilarity_search_with_embedding_id_by_vector(
- embedding=query_embedding,
- k=k_per_link or 10,
- filter=metadata_filter,
- )
- )
-
- results = await asyncio.gather(*tasks)
-
- for result in results:
- for doc, embedding, doc_id in result:
- if doc_id not in retrieved_docs:
- retrieved_docs[doc_id] = doc
- if doc_id not in targets:
- node = _doc_to_node(doc=doc)
- targets[doc_id] = AdjacentNode(node=node, embedding=embedding)
-
- # TODO: Consider a combined limit based on the similarity and/or
- # predicated MMR score?
- return targets.values()
-
- @staticmethod
- def _build_docs_from_texts(
- texts: List[str],
- metadatas: Optional[List[dict]] = None,
- ids: Optional[List[str]] = None,
- ) -> List[Document]:
- docs: List[Document] = []
- for i, text in enumerate(texts):
- doc = Document(
- page_content=text,
- )
- if metadatas is not None:
- doc.metadata = metadatas[i]
- if ids is not None:
- doc.id = ids[i]
- docs.append(doc)
- return docs
-
- @classmethod
- def from_texts(
- cls: Type[CGVST],
- texts: List[str],
- embedding: Embeddings,
- metadatas: Optional[List[dict]] = None,
- *,
- session: Optional[Session] = None,
- keyspace: Optional[str] = None,
- table_name: str = "",
- ids: Optional[List[str]] = None,
- ttl_seconds: Optional[int] = None,
- body_index_options: Optional[List[Tuple[str, Any]]] = None,
- metadata_deny_list: Optional[list[str]] = None,
- **kwargs: Any,
- ) -> CGVST:
- """Create a CassandraGraphVectorStore from raw texts.
-
- Args:
- texts: Texts to add to the vectorstore.
- embedding: Embedding function to use.
- metadatas: Optional list of metadatas associated with the texts.
- session: Cassandra driver session.
- If not provided, it is resolved from cassio.
- keyspace: Cassandra key space.
- If not provided, it is resolved from cassio.
- table_name: Cassandra table (required).
- ids: Optional list of IDs associated with the texts.
- ttl_seconds: Optional time-to-live for the added texts.
- body_index_options: Optional options used to create the body index.
- Eg. body_index_options = [cassio.table.cql.STANDARD_ANALYZER]
- metadata_deny_list: Optional list of metadata keys to not index.
- i.e. to fine-tune which of the metadata fields are indexed.
- Note: if you plan to have massive unique text metadata entries,
- consider not indexing them for performance
- (and to overcome max-length limitations).
- Note: the `metadata_indexing` parameter from
- langchain_community.utilities.cassandra.Cassandra is not
- exposed since CassandraGraphVectorStore only supports the
- deny_list option.
-
- Returns:
- a CassandraGraphVectorStore.
- """
- docs = cls._build_docs_from_texts(
- texts=texts,
- metadatas=metadatas,
- ids=ids,
- )
-
- return cls.from_documents(
- documents=docs,
- embedding=embedding,
- session=session,
- keyspace=keyspace,
- table_name=table_name,
- ttl_seconds=ttl_seconds,
- body_index_options=body_index_options,
- metadata_deny_list=metadata_deny_list,
- **kwargs,
- )
-
- @classmethod
- async def afrom_texts(
- cls: Type[CGVST],
- texts: List[str],
- embedding: Embeddings,
- metadatas: Optional[List[dict]] = None,
- *,
- session: Optional[Session] = None,
- keyspace: Optional[str] = None,
- table_name: str = "",
- ids: Optional[List[str]] = None,
- ttl_seconds: Optional[int] = None,
- body_index_options: Optional[List[Tuple[str, Any]]] = None,
- metadata_deny_list: Optional[list[str]] = None,
- **kwargs: Any,
- ) -> CGVST:
- """Create a CassandraGraphVectorStore from raw texts.
-
- Args:
- texts: Texts to add to the vectorstore.
- embedding: Embedding function to use.
- metadatas: Optional list of metadatas associated with the texts.
- session: Cassandra driver session.
- If not provided, it is resolved from cassio.
- keyspace: Cassandra key space.
- If not provided, it is resolved from cassio.
- table_name: Cassandra table (required).
- ids: Optional list of IDs associated with the texts.
- ttl_seconds: Optional time-to-live for the added texts.
- body_index_options: Optional options used to create the body index.
- Eg. body_index_options = [cassio.table.cql.STANDARD_ANALYZER]
- metadata_deny_list: Optional list of metadata keys to not index.
- i.e. to fine-tune which of the metadata fields are indexed.
- Note: if you plan to have massive unique text metadata entries,
- consider not indexing them for performance
- (and to overcome max-length limitations).
- Note: the `metadata_indexing` parameter from
- langchain_community.utilities.cassandra.Cassandra is not
- exposed since CassandraGraphVectorStore only supports the
- deny_list option.
-
- Returns:
- a CassandraGraphVectorStore.
- """
- docs = cls._build_docs_from_texts(
- texts=texts,
- metadatas=metadatas,
- ids=ids,
- )
-
- return await cls.afrom_documents(
- documents=docs,
- embedding=embedding,
- session=session,
- keyspace=keyspace,
- table_name=table_name,
- ttl_seconds=ttl_seconds,
- body_index_options=body_index_options,
- metadata_deny_list=metadata_deny_list,
- **kwargs,
- )
-
- @staticmethod
- def _add_ids_to_docs(
- docs: List[Document],
- ids: Optional[List[str]] = None,
- ) -> List[Document]:
- if ids is not None:
- for doc, doc_id in zip(docs, ids):
- doc.id = doc_id
- return docs
-
- @classmethod
- def from_documents(
- cls: Type[CGVST],
- documents: List[Document],
- embedding: Embeddings,
- *,
- session: Optional[Session] = None,
- keyspace: Optional[str] = None,
- table_name: str = "",
- ids: Optional[List[str]] = None,
- ttl_seconds: Optional[int] = None,
- body_index_options: Optional[List[Tuple[str, Any]]] = None,
- metadata_deny_list: Optional[list[str]] = None,
- **kwargs: Any,
- ) -> CGVST:
- """Create a CassandraGraphVectorStore from a document list.
-
- Args:
- documents: Documents to add to the vectorstore.
- embedding: Embedding function to use.
- session: Cassandra driver session.
- If not provided, it is resolved from cassio.
- keyspace: Cassandra key space.
- If not provided, it is resolved from cassio.
- table_name: Cassandra table (required).
- ids: Optional list of IDs associated with the documents.
- ttl_seconds: Optional time-to-live for the added documents.
- body_index_options: Optional options used to create the body index.
- Eg. body_index_options = [cassio.table.cql.STANDARD_ANALYZER]
- metadata_deny_list: Optional list of metadata keys to not index.
- i.e. to fine-tune which of the metadata fields are indexed.
- Note: if you plan to have massive unique text metadata entries,
- consider not indexing them for performance
- (and to overcome max-length limitations).
- Note: the `metadata_indexing` parameter from
- langchain_community.utilities.cassandra.Cassandra is not
- exposed since CassandraGraphVectorStore only supports the
- deny_list option.
-
- Returns:
- a CassandraGraphVectorStore.
- """
- store = cls(
- embedding=embedding,
- session=session,
- keyspace=keyspace,
- table_name=table_name,
- ttl_seconds=ttl_seconds,
- body_index_options=body_index_options,
- metadata_deny_list=metadata_deny_list,
- **kwargs,
- )
- store.add_documents(documents=cls._add_ids_to_docs(docs=documents, ids=ids))
- return store
-
- @classmethod
- async def afrom_documents(
- cls: Type[CGVST],
- documents: List[Document],
- embedding: Embeddings,
- *,
- session: Optional[Session] = None,
- keyspace: Optional[str] = None,
- table_name: str = "",
- ids: Optional[List[str]] = None,
- ttl_seconds: Optional[int] = None,
- body_index_options: Optional[List[Tuple[str, Any]]] = None,
- metadata_deny_list: Optional[list[str]] = None,
- **kwargs: Any,
- ) -> CGVST:
- """Create a CassandraGraphVectorStore from a document list.
-
- Args:
- documents: Documents to add to the vectorstore.
- embedding: Embedding function to use.
- session: Cassandra driver session.
- If not provided, it is resolved from cassio.
- keyspace: Cassandra key space.
- If not provided, it is resolved from cassio.
- table_name: Cassandra table (required).
- ids: Optional list of IDs associated with the documents.
- ttl_seconds: Optional time-to-live for the added documents.
- body_index_options: Optional options used to create the body index.
- Eg. body_index_options = [cassio.table.cql.STANDARD_ANALYZER]
- metadata_deny_list: Optional list of metadata keys to not index.
- i.e. to fine-tune which of the metadata fields are indexed.
- Note: if you plan to have massive unique text metadata entries,
- consider not indexing them for performance
- (and to overcome max-length limitations).
- Note: the `metadata_indexing` parameter from
- langchain_community.utilities.cassandra.Cassandra is not
- exposed since CassandraGraphVectorStore only supports the
- deny_list option.
-
-
- Returns:
- a CassandraGraphVectorStore.
- """
- store = cls(
- embedding=embedding,
- session=session,
- keyspace=keyspace,
- table_name=table_name,
- ttl_seconds=ttl_seconds,
- setup_mode=SetupMode.ASYNC,
- body_index_options=body_index_options,
- metadata_deny_list=metadata_deny_list,
- **kwargs,
- )
- await store.aadd_documents(
- documents=cls._add_ids_to_docs(docs=documents, ids=ids)
- )
- return store
diff --git a/libs/community/langchain_community/graph_vectorstores/extractors/__init__.py b/libs/community/langchain_community/graph_vectorstores/extractors/__init__.py
deleted file mode 100644
index 8d6a829ef6..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/extractors/__init__.py
+++ /dev/null
@@ -1,41 +0,0 @@
-from langchain_community.graph_vectorstores.extractors.gliner_link_extractor import (
- GLiNERInput,
- GLiNERLinkExtractor,
-)
-from langchain_community.graph_vectorstores.extractors.hierarchy_link_extractor import (
- HierarchyInput,
- HierarchyLinkExtractor,
-)
-from langchain_community.graph_vectorstores.extractors.html_link_extractor import (
- HtmlInput,
- HtmlLinkExtractor,
-)
-from langchain_community.graph_vectorstores.extractors.keybert_link_extractor import (
- KeybertInput,
- KeybertLinkExtractor,
-)
-from langchain_community.graph_vectorstores.extractors.link_extractor import (
- LinkExtractor,
-)
-from langchain_community.graph_vectorstores.extractors.link_extractor_adapter import (
- LinkExtractorAdapter,
-)
-from langchain_community.graph_vectorstores.extractors.link_extractor_transformer import ( # noqa: E501
- LinkExtractorTransformer,
-)
-
-__all__ = [
- "GLiNERInput",
- "GLiNERLinkExtractor",
- "HierarchyInput",
- "HierarchyLinkExtractor",
- "HtmlInput",
- "HtmlLinkExtractor",
- "KeybertInput",
- "KeybertLinkExtractor",
- "LinkExtractor",
- "LinkExtractor",
- "LinkExtractorAdapter",
- "LinkExtractorAdapter",
- "LinkExtractorTransformer",
-]
diff --git a/libs/community/langchain_community/graph_vectorstores/extractors/gliner_link_extractor.py b/libs/community/langchain_community/graph_vectorstores/extractors/gliner_link_extractor.py
deleted file mode 100644
index f353ba4a1d..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/extractors/gliner_link_extractor.py
+++ /dev/null
@@ -1,166 +0,0 @@
-from typing import Any, Dict, Iterable, List, Optional, Set, Union
-
-from langchain_core._api import beta
-from langchain_core.documents import Document
-
-from langchain_community.graph_vectorstores.extractors.link_extractor import (
- LinkExtractor,
-)
-from langchain_community.graph_vectorstores.links import Link
-
-# TypeAlias is not available in Python 3.9, we can't use that or the newer `type`.
-GLiNERInput = Union[str, Document]
-
-
-@beta()
-class GLiNERLinkExtractor(LinkExtractor[GLiNERInput]):
- """Link documents with common named entities using `GLiNER`_.
-
- `GLiNER`_ is a Named Entity Recognition (NER) model capable of identifying any
- entity type using a bidirectional transformer encoder (BERT-like).
-
- The ``GLiNERLinkExtractor`` uses GLiNER to create links between documents that
- have named entities in common.
-
- Example::
-
- extractor = GLiNERLinkExtractor(
- labels=["Person", "Award", "Date", "Competitions", "Teams"]
- )
- results = extractor.extract_one("some long text...")
-
- .. _GLiNER: https://github.com/urchade/GLiNER
-
- .. seealso::
-
- - :mod:`How to use a graph vector store `
- - :class:`How to create links between documents `
-
- How to link Documents on common named entities
- ==============================================
-
- Preliminaries
- -------------
-
- Install the ``gliner`` package:
-
- .. code-block:: bash
-
- pip install -q langchain_community gliner
-
- Usage
- -----
-
- We load the ``state_of_the_union.txt`` file, chunk it, then for each chunk we
- extract named entity links and add them to the chunk.
-
- Using extract_one()
- ^^^^^^^^^^^^^^^^^^^
-
- We can use :meth:`extract_one` on a document to get the links and add the links
- to the document metadata with
- :meth:`~langchain_community.graph_vectorstores.links.add_links`::
-
- from langchain_community.document_loaders import TextLoader
- from langchain_community.graph_vectorstores import CassandraGraphVectorStore
- from langchain_community.graph_vectorstores.extractors import GLiNERLinkExtractor
- from langchain_community.graph_vectorstores.links import add_links
- from langchain_text_splitters import CharacterTextSplitter
-
- loader = TextLoader("state_of_the_union.txt")
- raw_documents = loader.load()
-
- text_splitter = CharacterTextSplitter(chunk_size=1000, chunk_overlap=0)
- documents = text_splitter.split_documents(raw_documents)
-
- ner_extractor = GLiNERLinkExtractor(["Person", "Topic"])
- for document in documents:
- links = ner_extractor.extract_one(document)
- add_links(document, links)
-
- print(documents[0].metadata)
-
- .. code-block:: output
-
- {'source': 'state_of_the_union.txt', 'links': [Link(kind='entity:Person', direction='bidir', tag='President Zelenskyy'), Link(kind='entity:Person', direction='bidir', tag='Vladimir Putin')]}
-
- Using LinkExtractorTransformer
- ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
-
- Using the :class:`~langchain_community.graph_vectorstores.extractors.link_extractor_transformer.LinkExtractorTransformer`,
- we can simplify the link extraction::
-
- from langchain_community.document_loaders import TextLoader
- from langchain_community.graph_vectorstores.extractors import (
- GLiNERLinkExtractor,
- LinkExtractorTransformer,
- )
- from langchain_text_splitters import CharacterTextSplitter
-
- loader = TextLoader("state_of_the_union.txt")
- raw_documents = loader.load()
-
- text_splitter = CharacterTextSplitter(chunk_size=1000, chunk_overlap=0)
- documents = text_splitter.split_documents(raw_documents)
-
- ner_extractor = GLiNERLinkExtractor(["Person", "Topic"])
- transformer = LinkExtractorTransformer([ner_extractor])
- documents = transformer.transform_documents(documents)
-
- print(documents[0].metadata)
-
- .. code-block:: output
-
- {'source': 'state_of_the_union.txt', 'links': [Link(kind='entity:Person', direction='bidir', tag='President Zelenskyy'), Link(kind='entity:Person', direction='bidir', tag='Vladimir Putin')]}
-
- The documents with named entity links can then be added to a :class:`~langchain_community.graph_vectorstores.base.GraphVectorStore`::
-
- from langchain_community.graph_vectorstores import CassandraGraphVectorStore
-
- store = CassandraGraphVectorStore.from_documents(documents=documents, embedding=...)
-
- Args:
- labels: List of kinds of entities to extract.
- kind: Kind of links to produce with this extractor.
- model: GLiNER model to use.
- extract_kwargs: Keyword arguments to pass to GLiNER.
- """ # noqa: E501
-
- def __init__(
- self,
- labels: List[str],
- *,
- kind: str = "entity",
- model: str = "urchade/gliner_mediumv2.1",
- extract_kwargs: Optional[Dict[str, Any]] = None,
- ):
- try:
- from gliner import GLiNER
-
- self._model = GLiNER.from_pretrained(model)
-
- except ImportError:
- raise ImportError(
- "gliner is required for GLiNERLinkExtractor. "
- "Please install it with `pip install gliner`."
- ) from None
-
- self._labels = labels
- self._kind = kind
- self._extract_kwargs = extract_kwargs or {}
-
- def extract_one(self, input: GLiNERInput) -> Set[Link]: # noqa: A002
- return next(iter(self.extract_many([input])))
-
- def extract_many(
- self,
- inputs: Iterable[GLiNERInput],
- ) -> Iterable[Set[Link]]:
- strs = [i if isinstance(i, str) else i.page_content for i in inputs]
- for entities in self._model.batch_predict_entities(
- strs, self._labels, **self._extract_kwargs
- ):
- yield {
- Link.bidir(kind=f"{self._kind}:{e['label']}", tag=e["text"])
- for e in entities
- }
diff --git a/libs/community/langchain_community/graph_vectorstores/extractors/hierarchy_link_extractor.py b/libs/community/langchain_community/graph_vectorstores/extractors/hierarchy_link_extractor.py
deleted file mode 100644
index d838210ade..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/extractors/hierarchy_link_extractor.py
+++ /dev/null
@@ -1,110 +0,0 @@
-from typing import Callable, List, Set
-
-from langchain_core._api import beta
-from langchain_core.documents import Document
-
-from langchain_community.graph_vectorstores.extractors.link_extractor import (
- LinkExtractor,
-)
-from langchain_community.graph_vectorstores.extractors.link_extractor_adapter import (
- LinkExtractorAdapter,
-)
-from langchain_community.graph_vectorstores.links import Link
-
-# TypeAlias is not available in Python 3.9, we can't use that or the newer `type`.
-HierarchyInput = List[str]
-
-_PARENT: str = "p:"
-_CHILD: str = "c:"
-_SIBLING: str = "s:"
-
-
-@beta()
-class HierarchyLinkExtractor(LinkExtractor[HierarchyInput]):
- def __init__(
- self,
- *,
- kind: str = "hierarchy",
- parent_links: bool = True,
- child_links: bool = False,
- sibling_links: bool = False,
- ):
- """Extract links from a document hierarchy.
-
- Example:
-
- .. code-block:: python
-
- # Given three paths (in this case, within the "Root" document):
- h1 = ["Root", "H1"]
- h1a = ["Root", "H1", "a"]
- h1b = ["Root", "H1", "b"]
-
- # Parent links `h1a` and `h1b` to `h1`.
- # Child links `h1` to `h1a` and `h1b`.
- # Sibling links `h1a` and `h1b` together (both directions).
-
- Example use with documents:
- .. code_block: python
- transformer = LinkExtractorTransformer([
- HierarchyLinkExtractor().as_document_extractor(
- # Assumes the "path" to each document is in the metadata.
- # Could split strings, etc.
- lambda doc: doc.metadata.get("path", [])
- )
- ])
- linked = transformer.transform_documents(docs)
-
- Args:
- kind: Kind of links to produce with this extractor.
- parent_links: Link from a section to its parent.
- child_links: Link from a section to its children.
- sibling_links: Link from a section to other sections with the same parent.
- """
- self._kind = kind
- self._parent_links = parent_links
- self._child_links = child_links
- self._sibling_links = sibling_links
-
- def as_document_extractor(
- self, hierarchy: Callable[[Document], HierarchyInput]
- ) -> LinkExtractor[Document]:
- """Create a LinkExtractor from `Document`.
-
- Args:
- hierarchy: Function that returns the path for the given document.
-
- Returns:
- A `LinkExtractor[Document]` suitable for application to `Documents` directly
- or with `LinkExtractorTransformer`.
- """
- return LinkExtractorAdapter(underlying=self, transform=hierarchy)
-
- def extract_one(
- self,
- input: HierarchyInput,
- ) -> Set[Link]:
- this_path = "/".join(input)
- parent_path = None
-
- links = set()
- if self._parent_links:
- # This is linked from everything with this parent path.
- links.add(Link.incoming(kind=self._kind, tag=_PARENT + this_path))
- if self._child_links:
- # This is linked to every child with this as it's "parent" path.
- links.add(Link.outgoing(kind=self._kind, tag=_CHILD + this_path))
-
- if len(input) >= 1:
- parent_path = "/".join(input[0:-1])
- if self._parent_links and len(input) > 1:
- # This is linked to the nodes with the given parent path.
- links.add(Link.outgoing(kind=self._kind, tag=_PARENT + parent_path))
- if self._child_links and len(input) > 1:
- # This is linked from every node with the given parent path.
- links.add(Link.incoming(kind=self._kind, tag=_CHILD + parent_path))
- if self._sibling_links:
- # This is a sibling of everything with the same parent.
- links.add(Link.bidir(kind=self._kind, tag=_SIBLING + parent_path))
-
- return links
diff --git a/libs/community/langchain_community/graph_vectorstores/extractors/html_link_extractor.py b/libs/community/langchain_community/graph_vectorstores/extractors/html_link_extractor.py
deleted file mode 100644
index 8aee9767d4..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/extractors/html_link_extractor.py
+++ /dev/null
@@ -1,291 +0,0 @@
-from __future__ import annotations
-
-from dataclasses import dataclass
-from typing import TYPE_CHECKING, List, Optional, Set, Union
-from urllib.parse import urldefrag, urljoin, urlparse
-
-from langchain_core._api import beta
-from langchain_core.documents import Document
-
-from langchain_community.graph_vectorstores import Link
-from langchain_community.graph_vectorstores.extractors.link_extractor import (
- LinkExtractor,
-)
-from langchain_community.graph_vectorstores.extractors.link_extractor_adapter import (
- LinkExtractorAdapter,
-)
-
-if TYPE_CHECKING:
- from bs4 import BeautifulSoup
- from bs4.element import Tag
-
-
-def _parse_url(link: Tag, page_url: str, drop_fragments: bool = True) -> Optional[str]:
- href = link.get("href")
- if href is None:
- return None
- url = urlparse(href)
- if url.scheme not in ["http", "https", ""]:
- return None
-
- # Join the HREF with the page_url to convert relative paths to absolute.
- url = str(urljoin(page_url, href))
-
- # Fragments would be useful if we chunked a page based on section.
- # Then, each chunk would have a different URL based on the fragment.
- # Since we aren't doing that yet, they just "break" links. So, drop
- # the fragment.
- if drop_fragments:
- return urldefrag(url).url
- return url
-
-
-def _parse_hrefs(
- soup: BeautifulSoup, url: str, drop_fragments: bool = True
-) -> Set[str]:
- soup_links: List[Tag] = soup.find_all("a")
- links: Set[str] = set()
-
- for link in soup_links:
- parse_url = _parse_url(link, page_url=url, drop_fragments=drop_fragments)
- # Remove self links and entries for any 'a' tag that failed to parse
- # (didn't have href, or invalid domain, etc.)
- if parse_url and parse_url != url:
- links.add(parse_url)
-
- return links
-
-
-@dataclass
-class HtmlInput:
- content: Union[str, BeautifulSoup]
- base_url: str
-
-
-@beta()
-class HtmlLinkExtractor(LinkExtractor[HtmlInput]):
- def __init__(self, *, kind: str = "hyperlink", drop_fragments: bool = True):
- """Extract hyperlinks from HTML content.
-
- Expects the input to be an HTML string or a `BeautifulSoup` object.
-
- Example::
-
- extractor = HtmlLinkExtractor()
- results = extractor.extract_one(HtmlInput(html, url))
-
- .. seealso::
-
- - :mod:`How to use a graph vector store `
- - :class:`How to create links between documents `
-
- How to link Documents on hyperlinks in HTML
- ===========================================
-
- Preliminaries
- -------------
-
- Install the ``beautifulsoup4`` package:
-
- .. code-block:: bash
-
- pip install -q langchain_community beautifulsoup4
-
- Usage
- -----
-
- For this example, we'll scrape 2 HTML pages that have an hyperlink from one
- page to the other using an ``AsyncHtmlLoader``.
- Then we use the ``HtmlLinkExtractor`` to create the links in the documents.
-
- Using extract_one()
- ^^^^^^^^^^^^^^^^^^^
-
- We can use :meth:`extract_one` on a document to get the links and add the links
- to the document metadata with
- :meth:`~langchain_community.graph_vectorstores.links.add_links`::
-
- from langchain_community.document_loaders import AsyncHtmlLoader
- from langchain_community.graph_vectorstores.extractors import (
- HtmlInput,
- HtmlLinkExtractor,
- )
- from langchain_community.graph_vectorstores.links import add_links
- from langchain_core.documents import Document
-
- loader = AsyncHtmlLoader(
- [
- "https://python.langchain.com/docs/integrations/providers/astradb/",
- "https://docs.datastax.com/en/astra/home/astra.html",
- ]
- )
-
- documents = loader.load()
-
- html_extractor = HtmlLinkExtractor()
-
- for doc in documents:
- links = html_extractor.extract_one(HtmlInput(doc.page_content, url))
- add_links(doc, links)
-
- documents[0].metadata["links"][:5]
-
- .. code-block:: output
-
- [Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/integrations/providers/spreedly/'),
- Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/integrations/providers/nvidia/'),
- Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/integrations/providers/ray_serve/'),
- Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/integrations/providers/bageldb/'),
- Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/introduction/')]
-
- Using as_document_extractor()
- ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
-
- If you use a document loader that returns the raw HTML and that sets the source
- key in the document metadata such as ``AsyncHtmlLoader``,
- you can simplify by using :meth:`as_document_extractor` that takes directly a
- ``Document`` as input::
-
- from langchain_community.document_loaders import AsyncHtmlLoader
- from langchain_community.graph_vectorstores.extractors import HtmlLinkExtractor
- from langchain_community.graph_vectorstores.links import add_links
-
- loader = AsyncHtmlLoader(
- [
- "https://python.langchain.com/docs/integrations/providers/astradb/",
- "https://docs.datastax.com/en/astra/home/astra.html",
- ]
- )
- documents = loader.load()
- html_extractor = HtmlLinkExtractor().as_document_extractor()
-
- for document in documents:
- links = html_extractor.extract_one(document)
- add_links(document, links)
-
- documents[0].metadata["links"][:5]
-
- .. code-block:: output
-
- [Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/integrations/providers/spreedly/'),
- Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/integrations/providers/nvidia/'),
- Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/integrations/providers/ray_serve/'),
- Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/integrations/providers/bageldb/'),
- Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/introduction/')]
-
- Using LinkExtractorTransformer
- ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
-
- Using the :class:`~langchain_community.graph_vectorstores.extractors.link_extractor_transformer.LinkExtractorTransformer`,
- we can simplify the link extraction::
-
- from langchain_community.document_loaders import AsyncHtmlLoader
- from langchain_community.graph_vectorstores.extractors import (
- HtmlLinkExtractor,
- LinkExtractorTransformer,
- )
- from langchain_community.graph_vectorstores.links import add_links
-
- loader = AsyncHtmlLoader(
- [
- "https://python.langchain.com/docs/integrations/providers/astradb/",
- "https://docs.datastax.com/en/astra/home/astra.html",
- ]
- )
-
- documents = loader.load()
- transformer = LinkExtractorTransformer([HtmlLinkExtractor().as_document_extractor()])
- documents = transformer.transform_documents(documents)
-
- documents[0].metadata["links"][:5]
-
- .. code-block:: output
-
- [Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/integrations/providers/spreedly/'),
- Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/integrations/providers/nvidia/'),
- Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/integrations/providers/ray_serve/'),
- Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/integrations/providers/bageldb/'),
- Link(kind='hyperlink', direction='out', tag='https://python.langchain.com/docs/introduction/')]
-
- We can check that there is a link from the first document to the second::
-
- for doc_to in documents:
- for link_to in doc_to.metadata["links"]:
- if link_to.direction == "in":
- for doc_from in documents:
- for link_from in doc_from.metadata["links"]:
- if (
- link_to.direction == "in"
- and link_from.direction == "out"
- and link_to.tag == link_from.tag
- ):
- print(
- f"Found link from {doc_from.metadata['source']} to {doc_to.metadata['source']}."
- )
-
- .. code-block:: output
-
- Found link from https://python.langchain.com/docs/integrations/providers/astradb/ to https://docs.datastax.com/en/astra/home/astra.html.
-
- The documents with URL links can then be added to a :class:`~langchain_community.graph_vectorstores.base.GraphVectorStore`::
-
- from langchain_community.graph_vectorstores import CassandraGraphVectorStore
-
- store = CassandraGraphVectorStore.from_documents(documents=documents, embedding=...)
-
- Args:
- kind: The kind of edge to extract. Defaults to ``hyperlink``.
- drop_fragments: Whether fragments in URLs and links should be
- dropped. Defaults to ``True``.
- """ # noqa: E501
- try:
- import bs4 # noqa:F401
- except ImportError as e:
- raise ImportError(
- "BeautifulSoup4 is required for HtmlLinkExtractor. "
- "Please install it with `pip install beautifulsoup4`."
- ) from e
-
- self._kind = kind
- self.drop_fragments = drop_fragments
-
- def as_document_extractor(
- self, url_metadata_key: str = "source"
- ) -> LinkExtractor[Document]:
- """Return a LinkExtractor that applies to documents.
-
- Note:
- Since the HtmlLinkExtractor parses HTML, if you use with other similar
- link extractors it may be more efficient to call the link extractors
- directly on the parsed BeautifulSoup object.
-
- Args:
- url_metadata_key: The name of the filed in document metadata with the URL of
- the document.
- """
- return LinkExtractorAdapter(
- underlying=self,
- transform=lambda doc: HtmlInput(
- doc.page_content, doc.metadata[url_metadata_key]
- ),
- )
-
- def extract_one(
- self,
- input: HtmlInput, # noqa: A002
- ) -> Set[Link]:
- content = input.content
- if isinstance(content, str):
- from bs4 import BeautifulSoup
-
- content = BeautifulSoup(content, "html.parser")
-
- base_url = input.base_url
- if self.drop_fragments:
- base_url = urldefrag(base_url).url
-
- hrefs = _parse_hrefs(content, base_url, self.drop_fragments)
-
- links = {Link.outgoing(kind=self._kind, tag=url) for url in hrefs}
- links.add(Link.incoming(kind=self._kind, tag=base_url))
- return links
diff --git a/libs/community/langchain_community/graph_vectorstores/extractors/keybert_link_extractor.py b/libs/community/langchain_community/graph_vectorstores/extractors/keybert_link_extractor.py
deleted file mode 100644
index 3844df84f7..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/extractors/keybert_link_extractor.py
+++ /dev/null
@@ -1,167 +0,0 @@
-from typing import Any, Dict, Iterable, Optional, Set, Union
-
-from langchain_core._api import beta
-from langchain_core.documents import Document
-
-from langchain_community.graph_vectorstores.extractors.link_extractor import (
- LinkExtractor,
-)
-from langchain_community.graph_vectorstores.links import Link
-
-KeybertInput = Union[str, Document]
-
-
-@beta()
-class KeybertLinkExtractor(LinkExtractor[KeybertInput]):
- def __init__(
- self,
- *,
- kind: str = "kw",
- embedding_model: str = "all-MiniLM-L6-v2",
- extract_keywords_kwargs: Optional[Dict[str, Any]] = None,
- ):
- """Extract keywords using `KeyBERT `_.
-
- KeyBERT is a minimal and easy-to-use keyword extraction technique that
- leverages BERT embeddings to create keywords and keyphrases that are most
- similar to a document.
-
- The KeybertLinkExtractor uses KeyBERT to create links between documents that
- have keywords in common.
-
- Example::
-
- extractor = KeybertLinkExtractor()
- results = extractor.extract_one("lorem ipsum...")
-
- .. seealso::
-
- - :mod:`How to use a graph vector store `
- - :class:`How to create links between documents `
-
- How to link Documents on common keywords using Keybert
- ======================================================
-
- Preliminaries
- -------------
-
- Install the keybert package:
-
- .. code-block:: bash
-
- pip install -q langchain_community keybert
-
- Usage
- -----
-
- We load the ``state_of_the_union.txt`` file, chunk it, then for each chunk we
- extract keyword links and add them to the chunk.
-
- Using extract_one()
- ^^^^^^^^^^^^^^^^^^^
-
- We can use :meth:`extract_one` on a document to get the links and add the links
- to the document metadata with
- :meth:`~langchain_community.graph_vectorstores.links.add_links`::
-
- from langchain_community.document_loaders import TextLoader
- from langchain_community.graph_vectorstores import CassandraGraphVectorStore
- from langchain_community.graph_vectorstores.extractors import KeybertLinkExtractor
- from langchain_community.graph_vectorstores.links import add_links
- from langchain_text_splitters import CharacterTextSplitter
-
- loader = TextLoader("state_of_the_union.txt")
-
- raw_documents = loader.load()
- text_splitter = CharacterTextSplitter(chunk_size=1000, chunk_overlap=0)
-
- documents = text_splitter.split_documents(raw_documents)
- keyword_extractor = KeybertLinkExtractor()
-
- for document in documents:
- links = keyword_extractor.extract_one(document)
- add_links(document, links)
-
- print(documents[0].metadata)
-
- .. code-block:: output
-
- {'source': 'state_of_the_union.txt', 'links': [Link(kind='kw', direction='bidir', tag='ukraine'), Link(kind='kw', direction='bidir', tag='ukrainian'), Link(kind='kw', direction='bidir', tag='putin'), Link(kind='kw', direction='bidir', tag='vladimir'), Link(kind='kw', direction='bidir', tag='russia')]}
-
- Using LinkExtractorTransformer
- ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
-
- Using the :class:`~langchain_community.graph_vectorstores.extractors.link_extractor_transformer.LinkExtractorTransformer`,
- we can simplify the link extraction::
-
- from langchain_community.document_loaders import TextLoader
- from langchain_community.graph_vectorstores.extractors import (
- KeybertLinkExtractor,
- LinkExtractorTransformer,
- )
- from langchain_text_splitters import CharacterTextSplitter
-
- loader = TextLoader("state_of_the_union.txt")
- raw_documents = loader.load()
-
- text_splitter = CharacterTextSplitter(chunk_size=1000, chunk_overlap=0)
- documents = text_splitter.split_documents(raw_documents)
-
- transformer = LinkExtractorTransformer([KeybertLinkExtractor()])
- documents = transformer.transform_documents(documents)
-
- print(documents[0].metadata)
-
- .. code-block:: output
-
- {'source': 'state_of_the_union.txt', 'links': [Link(kind='kw', direction='bidir', tag='ukraine'), Link(kind='kw', direction='bidir', tag='ukrainian'), Link(kind='kw', direction='bidir', tag='putin'), Link(kind='kw', direction='bidir', tag='vladimir'), Link(kind='kw', direction='bidir', tag='russia')]}
-
- The documents with keyword links can then be added to a :class:`~langchain_community.graph_vectorstores.base.GraphVectorStore`::
-
- from langchain_community.graph_vectorstores import CassandraGraphVectorStore
-
- store = CassandraGraphVectorStore.from_documents(documents=documents, embedding=...)
-
- Args:
- kind: Kind of links to produce with this extractor.
- embedding_model: Name of the embedding model to use with KeyBERT.
- extract_keywords_kwargs: Keyword arguments to pass to KeyBERT's
- ``extract_keywords`` method.
- """ # noqa: E501
- try:
- import keybert
-
- self._kw_model = keybert.KeyBERT(model=embedding_model)
- except ImportError:
- raise ImportError(
- "keybert is required for KeybertLinkExtractor. "
- "Please install it with `pip install keybert`."
- ) from None
-
- self._kind = kind
- self._extract_keywords_kwargs = extract_keywords_kwargs or {}
-
- def extract_one(self, input: KeybertInput) -> Set[Link]: # noqa: A002
- keywords = self._kw_model.extract_keywords(
- input if isinstance(input, str) else input.page_content,
- **self._extract_keywords_kwargs,
- )
- return {Link.bidir(kind=self._kind, tag=kw[0]) for kw in keywords}
-
- def extract_many(
- self,
- inputs: Iterable[KeybertInput],
- ) -> Iterable[Set[Link]]:
- inputs = list(inputs)
- if len(inputs) == 1:
- # Even though we pass a list, if it contains one item, keybert will
- # flatten it. This means it's easier to just call the special case
- # for one item.
- yield self.extract_one(inputs[0])
- elif len(inputs) > 1:
- strs = [i if isinstance(i, str) else i.page_content for i in inputs]
- extracted = self._kw_model.extract_keywords(
- strs, **self._extract_keywords_kwargs
- )
- for keywords in extracted:
- yield {Link.bidir(kind=self._kind, tag=kw[0]) for kw in keywords}
diff --git a/libs/community/langchain_community/graph_vectorstores/extractors/link_extractor.py b/libs/community/langchain_community/graph_vectorstores/extractors/link_extractor.py
deleted file mode 100644
index bb141dccc5..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/extractors/link_extractor.py
+++ /dev/null
@@ -1,39 +0,0 @@
-from __future__ import annotations
-
-from abc import ABC, abstractmethod
-from typing import Generic, Iterable, Set, TypeVar
-
-from langchain_core._api import beta
-
-from langchain_community.graph_vectorstores import Link
-
-InputT = TypeVar("InputT")
-
-METADATA_LINKS_KEY = "links"
-
-
-@beta()
-class LinkExtractor(ABC, Generic[InputT]):
- """Interface for extracting links (incoming, outgoing, bidirectional)."""
-
- @abstractmethod
- def extract_one(self, input: InputT) -> Set[Link]:
- """Add edges from each `input` to the corresponding documents.
-
- Args:
- input: The input content to extract edges from.
-
- Returns:
- Set of links extracted from the input.
- """
-
- def extract_many(self, inputs: Iterable[InputT]) -> Iterable[Set[Link]]:
- """Add edges from each `input` to the corresponding documents.
-
- Args:
- inputs: The input content to extract edges from.
-
- Returns:
- Iterable over the set of links extracted from the input.
- """
- return map(self.extract_one, inputs)
diff --git a/libs/community/langchain_community/graph_vectorstores/extractors/link_extractor_adapter.py b/libs/community/langchain_community/graph_vectorstores/extractors/link_extractor_adapter.py
deleted file mode 100644
index 73b6761eff..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/extractors/link_extractor_adapter.py
+++ /dev/null
@@ -1,29 +0,0 @@
-from typing import Callable, Iterable, Set, TypeVar
-
-from langchain_core._api import beta
-
-from langchain_community.graph_vectorstores import Link
-from langchain_community.graph_vectorstores.extractors.link_extractor import (
- LinkExtractor,
-)
-
-InputT = TypeVar("InputT")
-UnderlyingInputT = TypeVar("UnderlyingInputT")
-
-
-@beta()
-class LinkExtractorAdapter(LinkExtractor[InputT]):
- def __init__(
- self,
- underlying: LinkExtractor[UnderlyingInputT],
- transform: Callable[[InputT], UnderlyingInputT],
- ) -> None:
- self._underlying = underlying
- self._transform = transform
-
- def extract_one(self, input: InputT) -> Set[Link]: # noqa: A002
- return self._underlying.extract_one(self._transform(input))
-
- def extract_many(self, inputs: Iterable[InputT]) -> Iterable[Set[Link]]:
- underlying_inputs = map(self._transform, inputs)
- return self._underlying.extract_many(underlying_inputs)
diff --git a/libs/community/langchain_community/graph_vectorstores/extractors/link_extractor_transformer.py b/libs/community/langchain_community/graph_vectorstores/extractors/link_extractor_transformer.py
deleted file mode 100644
index ba78fa2100..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/extractors/link_extractor_transformer.py
+++ /dev/null
@@ -1,45 +0,0 @@
-from typing import Any, Sequence
-
-from langchain_core._api import beta
-from langchain_core.documents import Document
-from langchain_core.documents.transformers import BaseDocumentTransformer
-
-from langchain_community.graph_vectorstores.extractors.link_extractor import (
- LinkExtractor,
-)
-from langchain_community.graph_vectorstores.links import copy_with_links
-
-
-@beta()
-class LinkExtractorTransformer(BaseDocumentTransformer):
- """DocumentTransformer for applying one or more LinkExtractors.
-
- Example:
- .. code-block:: python
-
- extract_links = LinkExtractorTransformer([
- HtmlLinkExtractor().as_document_extractor(),
- ])
- extract_links.transform_documents(docs)
- """
-
- def __init__(self, link_extractors: Sequence[LinkExtractor[Document]]):
- """Create a DocumentTransformer which adds extracted links to each document."""
- self.link_extractors = link_extractors
-
- def transform_documents(
- self, documents: Sequence[Document], **kwargs: Any
- ) -> Sequence[Document]:
- # Implement `transform_docments` directly, so that LinkExtractors which operate
- # better in batch (`extract_many`) get a chance to do so.
-
- # Run each extractor over all documents.
- links_per_extractor = [e.extract_many(documents) for e in self.link_extractors]
-
- # Transpose the list of lists to pair each document with the tuple of links.
- links_per_document = zip(*links_per_extractor)
-
- return [
- copy_with_links(document, *links)
- for document, links in zip(documents, links_per_document)
- ]
diff --git a/libs/community/langchain_community/graph_vectorstores/links.py b/libs/community/langchain_community/graph_vectorstores/links.py
deleted file mode 100644
index 8f32b03d2f..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/links.py
+++ /dev/null
@@ -1,220 +0,0 @@
-from collections.abc import Iterable
-from dataclasses import dataclass
-from typing import Literal, Union
-
-from langchain_core._api import beta
-from langchain_core.documents import Document
-
-
-@beta()
-@dataclass(frozen=True)
-class Link:
- """A link to/from a tag of a given kind.
-
- Documents in a :class:`graph vector store `
- are connected via "links".
- Links form a bipartite graph between documents and tags: documents are connected
- to tags, and tags are connected to other documents.
- When documents are retrieved from a graph vector store, a pair of documents are
- connected with a depth of one if both documents are connected to the same tag.
-
- Links have a ``kind`` property, used to namespace different tag identifiers.
- For example a link to a keyword might use kind ``kw``, while a link to a URL might
- use kind ``url``.
- This allows the same tag value to be used in different contexts without causing
- name collisions.
-
- Links are directed. The directionality of links controls how the graph is
- traversed at retrieval time.
- For example, given documents ``A`` and ``B``, connected by links to tag ``T``:
-
- +----------+----------+---------------------------------+
- | A to T | B to T | Result |
- +==========+==========+=================================+
- | outgoing | incoming | Retrieval traverses from A to B |
- +----------+----------+---------------------------------+
- | incoming | incoming | No traversal from A to B |
- +----------+----------+---------------------------------+
- | outgoing | incoming | No traversal from A to B |
- +----------+----------+---------------------------------+
- | bidir | incoming | Retrieval traverses from A to B |
- +----------+----------+---------------------------------+
- | bidir | outgoing | No traversal from A to B |
- +----------+----------+---------------------------------+
- | outgoing | bidir | Retrieval traverses from A to B |
- +----------+----------+---------------------------------+
- | incoming | bidir | No traversal from A to B |
- +----------+----------+---------------------------------+
-
- Directed links make it possible to describe relationships such as term
- references / definitions: term definitions are generally relevant to any documents
- that use the term, but the full set of documents using a term generally aren't
- relevant to the term's definition.
-
- .. seealso::
-
- - :mod:`How to use a graph vector store `
- - :class:`How to link Documents on hyperlinks in HTML `
- - :class:`How to link Documents on common keywords (using KeyBERT) `
- - :class:`How to link Documents on common named entities (using GliNER) `
-
- How to add links to a Document
- ==============================
-
- How to create links
- -------------------
-
- You can create links using the Link class's constructors :meth:`incoming`,
- :meth:`outgoing`, and :meth:`bidir`::
-
- from langchain_community.graph_vectorstores.links import Link
-
- print(Link.bidir(kind="location", tag="Paris"))
-
- .. code-block:: output
-
- Link(kind='location', direction='bidir', tag='Paris')
-
- Extending documents with links
- ------------------------------
-
- Now that we know how to create links, let's associate them with some documents.
- These edges will strengthen the connection between documents that share a keyword
- when using a graph vector store to retrieve documents.
-
- First, we'll load some text and chunk it into smaller pieces.
- Then we'll add a link to each document to link them all together::
-
- from langchain_community.document_loaders import TextLoader
- from langchain_community.graph_vectorstores.links import add_links
- from langchain_text_splitters import CharacterTextSplitter
-
- loader = TextLoader("state_of_the_union.txt")
-
- raw_documents = loader.load()
- text_splitter = CharacterTextSplitter(chunk_size=1000, chunk_overlap=0)
- documents = text_splitter.split_documents(raw_documents)
-
- for doc in documents:
- add_links(doc, Link.bidir(kind="genre", tag="oratory"))
-
- print(documents[0].metadata)
-
- .. code-block:: output
-
- {'source': 'state_of_the_union.txt', 'links': [Link(kind='genre', direction='bidir', tag='oratory')]}
-
- As we can see, each document's metadata now includes a bidirectional link to the
- genre ``oratory``.
-
- The documents can then be added to a graph vector store::
-
- from langchain_community.graph_vectorstores import CassandraGraphVectorStore
-
- graph_vectorstore = CassandraGraphVectorStore.from_documents(
- documents=documents, embeddings=...
- )
-
- """ # noqa: E501
-
- kind: str
- """The kind of link. Allows different extractors to use the same tag name without
- creating collisions between extractors. For example “keyword” vs “url”."""
- direction: Literal["in", "out", "bidir"]
- """The direction of the link."""
- tag: str
- """The tag of the link."""
-
- @staticmethod
- def incoming(kind: str, tag: str) -> "Link":
- """Create an incoming link.
-
- Args:
- kind: the link kind.
- tag: the link tag.
- """
- return Link(kind=kind, direction="in", tag=tag)
-
- @staticmethod
- def outgoing(kind: str, tag: str) -> "Link":
- """Create an outgoing link.
-
- Args:
- kind: the link kind.
- tag: the link tag.
- """
- return Link(kind=kind, direction="out", tag=tag)
-
- @staticmethod
- def bidir(kind: str, tag: str) -> "Link":
- """Create a bidirectional link.
-
- Args:
- kind: the link kind.
- tag: the link tag.
- """
- return Link(kind=kind, direction="bidir", tag=tag)
-
-
-METADATA_LINKS_KEY = "links"
-
-
-@beta()
-def get_links(doc: Document) -> list[Link]:
- """Get the links from a document.
-
- Args:
- doc: The document to get the link tags from.
- Returns:
- The set of link tags from the document.
- """
-
- links = doc.metadata.setdefault(METADATA_LINKS_KEY, [])
- if not isinstance(links, list):
- # Convert to a list and remember that.
- links = list(links)
- doc.metadata[METADATA_LINKS_KEY] = links
- return links
-
-
-@beta()
-def add_links(doc: Document, *links: Union[Link, Iterable[Link]]) -> None:
- """Add links to the given metadata.
-
- Args:
- doc: The document to add the links to.
- *links: The links to add to the document.
- """
- links_in_metadata = get_links(doc)
- for link in links:
- if isinstance(link, Iterable):
- links_in_metadata.extend(link)
- else:
- links_in_metadata.append(link)
-
-
-@beta()
-def copy_with_links(doc: Document, *links: Union[Link, Iterable[Link]]) -> Document:
- """Return a document with the given links added.
-
- Args:
- doc: The document to add the links to.
- *links: The links to add to the document.
-
- Returns:
- A document with a shallow-copy of the metadata with the links added.
- """
- new_links = set(get_links(doc))
- for link in links:
- if isinstance(link, Iterable):
- new_links.update(link)
- else:
- new_links.add(link)
-
- return Document(
- page_content=doc.page_content,
- metadata={
- **doc.metadata,
- METADATA_LINKS_KEY: list(new_links),
- },
- )
diff --git a/libs/community/langchain_community/graph_vectorstores/mmr_helper.py b/libs/community/langchain_community/graph_vectorstores/mmr_helper.py
deleted file mode 100644
index 43aa8c0949..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/mmr_helper.py
+++ /dev/null
@@ -1,272 +0,0 @@
-"""Tools for the Graph Traversal Maximal Marginal Relevance (MMR) reranking."""
-
-from __future__ import annotations
-
-import dataclasses
-from typing import TYPE_CHECKING, Iterable
-
-import numpy as np
-
-from langchain_community.utils.math import cosine_similarity
-
-if TYPE_CHECKING:
- from numpy.typing import NDArray
-
-
-def _emb_to_ndarray(embedding: list[float]) -> NDArray[np.float32]:
- emb_array = np.array(embedding, dtype=np.float32)
- if emb_array.ndim == 1:
- emb_array = np.expand_dims(emb_array, axis=0)
- return emb_array
-
-
-NEG_INF = float("-inf")
-
-
-@dataclasses.dataclass
-class _Candidate:
- id: str
- similarity: float
- weighted_similarity: float
- weighted_redundancy: float
- score: float = dataclasses.field(init=False)
-
- def __post_init__(self) -> None:
- self.score = self.weighted_similarity - self.weighted_redundancy
-
- def update_redundancy(self, new_weighted_redundancy: float) -> None:
- if new_weighted_redundancy > self.weighted_redundancy:
- self.weighted_redundancy = new_weighted_redundancy
- self.score = self.weighted_similarity - self.weighted_redundancy
-
-
-class MmrHelper:
- """Helper for executing an MMR traversal query.
-
- Args:
- query_embedding: The embedding of the query to use for scoring.
- lambda_mult: Number between 0 and 1 that determines the degree
- of diversity among the results with 0 corresponding to maximum
- diversity and 1 to minimum diversity. Defaults to 0.5.
- score_threshold: Only documents with a score greater than or equal
- this threshold will be chosen. Defaults to -infinity.
- """
-
- dimensions: int
- """Dimensions of the embedding."""
-
- query_embedding: NDArray[np.float32]
- """Embedding of the query as a (1,dim) ndarray."""
-
- lambda_mult: float
- """Number between 0 and 1.
-
- Determines the degree of diversity among the results with 0 corresponding to
- maximum diversity and 1 to minimum diversity."""
-
- lambda_mult_complement: float
- """1 - lambda_mult."""
-
- score_threshold: float
- """Only documents with a score greater than or equal to this will be chosen."""
-
- selected_ids: list[str]
- """List of selected IDs (in selection order)."""
-
- selected_mmr_scores: list[float]
- """List of MMR score at the time each document is selected."""
-
- selected_similarity_scores: list[float]
- """List of similarity score for each selected document."""
-
- selected_embeddings: NDArray[np.float32]
- """(N, dim) ndarray with a row for each selected node."""
-
- candidate_id_to_index: dict[str, int]
- """Dictionary of candidate IDs to indices in candidates and candidate_embeddings."""
- candidates: list[_Candidate]
- """List containing information about candidates.
-
- Same order as rows in `candidate_embeddings`.
- """
- candidate_embeddings: NDArray[np.float32]
- """(N, dim) ndarray with a row for each candidate."""
-
- best_score: float
- best_id: str | None
-
- def __init__(
- self,
- k: int,
- query_embedding: list[float],
- lambda_mult: float = 0.5,
- score_threshold: float = NEG_INF,
- ) -> None:
- """Create a new Traversal MMR helper."""
- self.query_embedding = _emb_to_ndarray(query_embedding)
- self.dimensions = self.query_embedding.shape[1]
-
- self.lambda_mult = lambda_mult
- self.lambda_mult_complement = 1 - lambda_mult
- self.score_threshold = score_threshold
-
- self.selected_ids = []
- self.selected_similarity_scores = []
- self.selected_mmr_scores = []
-
- # List of selected embeddings (in selection order).
- self.selected_embeddings = np.ndarray((k, self.dimensions), dtype=np.float32)
-
- self.candidate_id_to_index = {}
-
- # List of the candidates.
- self.candidates = []
- # numpy n-dimensional array of the candidate embeddings.
- self.candidate_embeddings = np.ndarray((0, self.dimensions), dtype=np.float32)
-
- self.best_score = NEG_INF
- self.best_id = None
-
- def candidate_ids(self) -> Iterable[str]:
- """Return the IDs of the candidates."""
- return self.candidate_id_to_index.keys()
-
- def _already_selected_embeddings(self) -> NDArray[np.float32]:
- """Return the selected embeddings sliced to the already assigned values."""
- selected = len(self.selected_ids)
- return np.vsplit(self.selected_embeddings, [selected])[0]
-
- def _pop_candidate(self, candidate_id: str) -> tuple[float, NDArray[np.float32]]:
- """Pop the candidate with the given ID.
-
- Returns:
- The similarity score and embedding of the candidate.
- """
- # Get the embedding for the id.
- index = self.candidate_id_to_index.pop(candidate_id)
- if self.candidates[index].id != candidate_id:
- msg = (
- "ID in self.candidate_id_to_index doesn't match the ID of the "
- "corresponding index in self.candidates"
- )
- raise ValueError(msg)
- embedding: NDArray[np.float32] = self.candidate_embeddings[index].copy()
-
- # Swap that index with the last index in the candidates and
- # candidate_embeddings.
- last_index = self.candidate_embeddings.shape[0] - 1
-
- similarity = 0.0
- if index == last_index:
- # Already the last item. We don't need to swap.
- similarity = self.candidates.pop().similarity
- else:
- self.candidate_embeddings[index] = self.candidate_embeddings[last_index]
-
- similarity = self.candidates[index].similarity
-
- old_last = self.candidates.pop()
- self.candidates[index] = old_last
- self.candidate_id_to_index[old_last.id] = index
-
- self.candidate_embeddings = np.vsplit(self.candidate_embeddings, [last_index])[
- 0
- ]
-
- return similarity, embedding
-
- def pop_best(self) -> str | None:
- """Select and pop the best item being considered.
-
- Updates the consideration set based on it.
-
- Returns:
- A tuple containing the ID of the best item.
- """
- if self.best_id is None or self.best_score < self.score_threshold:
- return None
-
- # Get the selection and remove from candidates.
- selected_id = self.best_id
- selected_similarity, selected_embedding = self._pop_candidate(selected_id)
-
- # Add the ID and embedding to the selected information.
- selection_index = len(self.selected_ids)
- self.selected_ids.append(selected_id)
- self.selected_mmr_scores.append(self.best_score)
- self.selected_similarity_scores.append(selected_similarity)
- self.selected_embeddings[selection_index] = selected_embedding
-
- # Reset the best score / best ID.
- self.best_score = NEG_INF
- self.best_id = None
-
- # Update the candidates redundancy, tracking the best node.
- if self.candidate_embeddings.shape[0] > 0:
- similarity = cosine_similarity(
- self.candidate_embeddings, np.expand_dims(selected_embedding, axis=0)
- )
- for index, candidate in enumerate(self.candidates):
- candidate.update_redundancy(similarity[index][0])
- if candidate.score > self.best_score:
- self.best_score = candidate.score
- self.best_id = candidate.id
-
- return selected_id
-
- def add_candidates(self, candidates: dict[str, list[float]]) -> None:
- """Add candidates to the consideration set."""
- # Determine the keys to actually include.
- # These are the candidates that aren't already selected
- # or under consideration.
- include_ids_set = set(candidates.keys())
- include_ids_set.difference_update(self.selected_ids)
- include_ids_set.difference_update(self.candidate_id_to_index.keys())
- include_ids = list(include_ids_set)
-
- # Now, build up a matrix of the remaining candidate embeddings.
- # And add them to the
- new_embeddings: NDArray[np.float32] = np.ndarray(
- (
- len(include_ids),
- self.dimensions,
- )
- )
- offset = self.candidate_embeddings.shape[0]
- for index, candidate_id in enumerate(include_ids):
- if candidate_id in include_ids:
- self.candidate_id_to_index[candidate_id] = offset + index
- embedding = candidates[candidate_id]
- new_embeddings[index] = embedding
-
- # Compute the similarity to the query.
- similarity = cosine_similarity(new_embeddings, self.query_embedding)
-
- # Compute the distance metrics of all of pairs in the selected set with
- # the new candidates.
- redundancy = cosine_similarity(
- new_embeddings, self._already_selected_embeddings()
- )
- for index, candidate_id in enumerate(include_ids):
- max_redundancy = 0.0
- if redundancy.shape[0] > 0:
- max_redundancy = redundancy[index].max()
- candidate = _Candidate(
- id=candidate_id,
- similarity=similarity[index][0],
- weighted_similarity=self.lambda_mult * similarity[index][0],
- weighted_redundancy=self.lambda_mult_complement * max_redundancy,
- )
- self.candidates.append(candidate)
-
- if candidate.score >= self.best_score:
- self.best_score = candidate.score
- self.best_id = candidate.id
-
- # Add the new embeddings to the candidate set.
- self.candidate_embeddings = np.vstack(
- (
- self.candidate_embeddings,
- new_embeddings,
- )
- )
diff --git a/libs/community/langchain_community/graph_vectorstores/networkx.py b/libs/community/langchain_community/graph_vectorstores/networkx.py
deleted file mode 100644
index 7a3c920297..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/networkx.py
+++ /dev/null
@@ -1,84 +0,0 @@
-"""Utilities for using Graph Vector Stores with networkx."""
-
-import typing
-
-from langchain_core.documents import Document
-
-from langchain_community.graph_vectorstores.links import get_links
-
-if typing.TYPE_CHECKING:
- import networkx as nx
-
-
-def documents_to_networkx(
- documents: typing.Iterable[Document],
- *,
- tag_nodes: bool = True,
-) -> "nx.DiGraph":
- """Return the networkx directed graph corresponding to the documents.
-
- Args:
- documents: The documents to convenrt to networkx.
- tag_nodes: If `True`, each tag will be rendered as a node, with edges
- to/from the corresponding documents. If `False`, edges will be
- between documents, with a label corresponding to the tag(s)
- connecting them.
- """
- import networkx as nx
-
- graph = nx.DiGraph()
-
- tag_ids: typing.Dict[typing.Tuple[str, str], str] = {}
- tag_labels: typing.Dict[str, str] = {}
- documents_by_incoming: typing.Dict[str, typing.Set[str]] = {}
-
- # First pass:
- # - Register tag IDs for each unique (kind, tag).
- # - If rendering tag nodes, add them to the graph.
- # - If not rendering tag nodes, create a dictionary of documents by incoming tags.
- for document in documents:
- if document.id is None:
- raise ValueError(f"Illegal graph document without ID: {document}")
-
- for link in get_links(document):
- tag_key = (link.kind, link.tag)
- tag_id = tag_ids.get(tag_key)
- if tag_id is None:
- tag_id = f"tag_{len(tag_ids)}"
- tag_ids[tag_key] = tag_id
-
- if tag_nodes:
- graph.add_node(tag_id, label=f"{link.kind}:{link.tag}")
-
- if not tag_nodes and (link.direction == "in" or link.direction == "bidir"):
- tag_labels[tag_id] = f"{link.kind}:{link.tag}"
- documents_by_incoming.setdefault(tag_id, set()).add(document.id)
-
- # Second pass:
- # - Render document nodes
- # - If rendering tag nodes, render edges to/from documents and tag nodes.
- # - If not rendering tag nodes, render edges to/from documents based on tags.
- for document in documents:
- graph.add_node(document.id, text=document.page_content)
-
- targets: typing.Dict[str, typing.List[str]] = {}
- for link in get_links(document):
- tag_id = tag_ids[(link.kind, link.tag)]
- if tag_nodes:
- if link.direction == "in" or link.direction == "bidir":
- graph.add_edge(tag_id, document.id)
- if link.direction == "out" or link.direction == "bidir":
- graph.add_edge(document.id, tag_id)
- else:
- if link.direction == "out" or link.direction == "bidir":
- label = tag_labels[tag_id]
- for target in documents_by_incoming[tag_id]:
- if target != document.id:
- targets.setdefault(target, []).append(label)
-
- # Avoid a multigraph by collecting the list of labels for each edge.
- if not tag_nodes:
- for target, labels in targets.items():
- graph.add_edge(document.id, target, label=str(labels))
-
- return graph
diff --git a/libs/community/langchain_community/graph_vectorstores/visualize.py b/libs/community/langchain_community/graph_vectorstores/visualize.py
deleted file mode 100644
index 8c745a10d3..0000000000
--- a/libs/community/langchain_community/graph_vectorstores/visualize.py
+++ /dev/null
@@ -1,122 +0,0 @@
-import re
-from typing import TYPE_CHECKING, Dict, Iterable, Optional, Tuple
-
-from langchain_core._api import beta
-from langchain_core.documents import Document
-
-from langchain_community.graph_vectorstores.links import get_links
-
-if TYPE_CHECKING:
- import graphviz
-
-
-def _escape_id(id: str) -> str:
- return id.replace(":", "_")
-
-
-_EDGE_DIRECTION = {
- "in": "back",
- "out": "forward",
- "bidir": "both",
-}
-
-_WORD_RE = re.compile(r"\s*\S+")
-
-
-def _split_prefix(s: str, max_chars: int = 50) -> str:
- words = _WORD_RE.finditer(s)
-
- split = min(len(s), max_chars)
- for word in words:
- if word.end(0) > max_chars:
- break
- split = word.end(0)
-
- if split == len(s):
- return s
- else:
- return f"{s[0:split]}..."
-
-
-@beta()
-def render_graphviz(
- documents: Iterable[Document],
- engine: Optional[str] = None,
- node_color: Optional[str] = None,
- node_colors: Optional[Dict[str, Optional[str]]] = None,
- skip_tags: Iterable[Tuple[str, str]] = (),
-) -> "graphviz.Digraph":
- """Render a collection of GraphVectorStore documents to GraphViz format.
-
- Args:
- documents: The documents to render.
- engine: GraphViz layout engine to use. `None` uses the default.
- node_color: Default node color.
- node_colors: Dictionary specifying colors of specific nodes. Useful for
- emphasizing nodes that were selected by MMR, or differ from other
- results.
- skip_tags: Set of tags to skip when rendering the graph. Specified as
- tuples containing the kind and tag.
-
- Returns:
- The "graphviz.Digraph" representing the nodes. May be printed to source,
- or rendered using `dot`.
-
- Note:
- To render the generated DOT source code, you also need to install Graphviz_
- (`download page `_,
- `archived versions `_,
- `installation procedure for Windows `_).
- """
- if node_colors is None:
- node_colors = {}
-
- try:
- import graphviz
- except (ImportError, ModuleNotFoundError):
- raise ImportError(
- "Could not import graphviz python package. "
- "Please install it with `pip install graphviz`."
- )
-
- graph = graphviz.Digraph(engine=engine)
- graph.attr(rankdir="LR")
- graph.attr("node", style="filled")
-
- skip_tags = set(skip_tags)
- tags: dict[Tuple[str, str], str] = {}
-
- for document in documents:
- id = document.id
- if id is None:
- raise ValueError(f"Illegal graph document without ID: {document}")
- escaped_id = _escape_id(id)
- color = node_colors[id] if id in node_colors else node_color
-
- node_label = "\n".join(
- [
- graphviz.escape(id),
- graphviz.escape(_split_prefix(document.page_content)),
- ]
- )
- graph.node(
- escaped_id,
- label=node_label,
- shape="note",
- fillcolor=color,
- tooltip=graphviz.escape(document.page_content),
- )
-
- for link in get_links(document):
- tag_key = (link.kind, link.tag)
- if tag_key in skip_tags:
- continue
-
- tag_id = tags.get(tag_key)
- if tag_id is None:
- tag_id = f"tag_{len(tags)}"
- tags[tag_key] = tag_id
- graph.node(tag_id, label=graphviz.escape(f"{link.kind}:{link.tag}"))
-
- graph.edge(escaped_id, tag_id, dir=_EDGE_DIRECTION[link.direction])
- return graph
diff --git a/libs/community/langchain_community/graphs/__init__.py b/libs/community/langchain_community/graphs/__init__.py
deleted file mode 100644
index 37bbf71b04..0000000000
--- a/libs/community/langchain_community/graphs/__init__.py
+++ /dev/null
@@ -1,95 +0,0 @@
-"""**Graphs** provide a natural language interface to graph databases."""
-
-import importlib
-from typing import TYPE_CHECKING, Any
-
-if TYPE_CHECKING:
- from langchain_community.graphs.arangodb_graph import (
- ArangoGraph,
- )
- from langchain_community.graphs.falkordb_graph import (
- FalkorDBGraph,
- )
- from langchain_community.graphs.gremlin_graph import (
- GremlinGraph,
- )
- from langchain_community.graphs.hugegraph import (
- HugeGraph,
- )
- from langchain_community.graphs.kuzu_graph import (
- KuzuGraph,
- )
- from langchain_community.graphs.memgraph_graph import (
- MemgraphGraph,
- )
- from langchain_community.graphs.nebula_graph import (
- NebulaGraph,
- )
- from langchain_community.graphs.neo4j_graph import (
- Neo4jGraph,
- )
- from langchain_community.graphs.neptune_graph import (
- BaseNeptuneGraph,
- NeptuneAnalyticsGraph,
- NeptuneGraph,
- )
- from langchain_community.graphs.neptune_rdf_graph import (
- NeptuneRdfGraph,
- )
- from langchain_community.graphs.networkx_graph import (
- NetworkxEntityGraph,
- )
- from langchain_community.graphs.ontotext_graphdb_graph import (
- OntotextGraphDBGraph,
- )
- from langchain_community.graphs.rdf_graph import (
- RdfGraph,
- )
- from langchain_community.graphs.tigergraph_graph import (
- TigerGraph,
- )
-
-__all__ = [
- "ArangoGraph",
- "FalkorDBGraph",
- "GremlinGraph",
- "HugeGraph",
- "KuzuGraph",
- "BaseNeptuneGraph",
- "MemgraphGraph",
- "NebulaGraph",
- "Neo4jGraph",
- "NeptuneGraph",
- "NeptuneRdfGraph",
- "NeptuneAnalyticsGraph",
- "NetworkxEntityGraph",
- "OntotextGraphDBGraph",
- "RdfGraph",
- "TigerGraph",
-]
-
-_module_lookup = {
- "ArangoGraph": "langchain_community.graphs.arangodb_graph",
- "FalkorDBGraph": "langchain_community.graphs.falkordb_graph",
- "GremlinGraph": "langchain_community.graphs.gremlin_graph",
- "HugeGraph": "langchain_community.graphs.hugegraph",
- "KuzuGraph": "langchain_community.graphs.kuzu_graph",
- "MemgraphGraph": "langchain_community.graphs.memgraph_graph",
- "NebulaGraph": "langchain_community.graphs.nebula_graph",
- "Neo4jGraph": "langchain_community.graphs.neo4j_graph",
- "BaseNeptuneGraph": "langchain_community.graphs.neptune_graph",
- "NeptuneAnalyticsGraph": "langchain_community.graphs.neptune_graph",
- "NeptuneGraph": "langchain_community.graphs.neptune_graph",
- "NeptuneRdfGraph": "langchain_community.graphs.neptune_rdf_graph",
- "NetworkxEntityGraph": "langchain_community.graphs.networkx_graph",
- "OntotextGraphDBGraph": "langchain_community.graphs.ontotext_graphdb_graph",
- "RdfGraph": "langchain_community.graphs.rdf_graph",
- "TigerGraph": "langchain_community.graphs.tigergraph_graph",
-}
-
-
-def __getattr__(name: str) -> Any:
- if name in _module_lookup:
- module = importlib.import_module(_module_lookup[name])
- return getattr(module, name)
- raise AttributeError(f"module {__name__} has no attribute {name}")
diff --git a/libs/community/langchain_community/graphs/age_graph.py b/libs/community/langchain_community/graphs/age_graph.py
deleted file mode 100644
index 116791ee5c..0000000000
--- a/libs/community/langchain_community/graphs/age_graph.py
+++ /dev/null
@@ -1,765 +0,0 @@
-from __future__ import annotations
-
-import json
-import re
-from hashlib import md5
-from typing import TYPE_CHECKING, Any, Dict, List, NamedTuple, Pattern, Tuple, Union
-
-from langchain_community.graphs.graph_document import GraphDocument
-from langchain_community.graphs.graph_store import GraphStore
-
-if TYPE_CHECKING:
- import psycopg2.extras
-
-
-class AGEQueryException(Exception):
- """Exception for the AGE queries."""
-
- def __init__(self, exception: Union[str, Dict]) -> None:
- if isinstance(exception, dict):
- self.message = exception["message"] if "message" in exception else "unknown"
- self.details = exception["details"] if "details" in exception else "unknown"
- else:
- self.message = exception
- self.details = "unknown"
-
- def get_message(self) -> str:
- return self.message
-
- def get_details(self) -> Any:
- return self.details
-
-
-class AGEGraph(GraphStore):
- """
- Apache AGE wrapper for graph operations.
-
- Args:
- graph_name (str): the name of the graph to connect to or create
- conf (Dict[str, Any]): the pgsql connection config passed directly
- to psycopg2.connect
- create (bool): if True and graph doesn't exist, attempt to create it
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- # python type mapping for providing readable types to LLM
- types = {
- "str": "STRING",
- "float": "DOUBLE",
- "int": "INTEGER",
- "list": "LIST",
- "dict": "MAP",
- "bool": "BOOLEAN",
- }
-
- # precompiled regex for checking chars in graph labels
- label_regex: Pattern = re.compile("[^0-9a-zA-Z]+")
-
- def __init__(
- self, graph_name: str, conf: Dict[str, Any], create: bool = True
- ) -> None:
- """Create a new AGEGraph instance."""
-
- self.graph_name = graph_name
-
- # check that psycopg2 is installed
- try:
- import psycopg2
- except ImportError:
- raise ImportError(
- "Could not import psycopg2 python package. "
- "Please install it with `pip install psycopg2`."
- )
-
- self.connection = psycopg2.connect(**conf)
-
- with self._get_cursor() as curs:
- # check if graph with name graph_name exists
- graph_id_query = (
- """SELECT graphid FROM ag_catalog.ag_graph WHERE name = '{}'""".format(
- graph_name
- )
- )
-
- curs.execute(graph_id_query)
- data = curs.fetchone()
-
- # if graph doesn't exist and create is True, create it
- if data is None:
- if create:
- create_statement = """
- SELECT ag_catalog.create_graph('{}');
- """.format(graph_name)
-
- try:
- curs.execute(create_statement)
- self.connection.commit()
- except psycopg2.Error as e:
- raise AGEQueryException(
- {
- "message": "Could not create the graph",
- "detail": str(e),
- }
- )
-
- else:
- raise Exception(
- (
- 'Graph "{}" does not exist in the database '
- + 'and "create" is set to False'
- ).format(graph_name)
- )
-
- curs.execute(graph_id_query)
- data = curs.fetchone()
-
- # store graph id and refresh the schema
- self.graphid = data.graphid
- self.refresh_schema()
-
- def _get_cursor(self) -> psycopg2.extras.NamedTupleCursor:
- """
- get cursor, load age extension and set search path
- """
-
- try:
- import psycopg2.extras
- except ImportError as e:
- raise ImportError(
- "Unable to import psycopg2, please install with "
- "`pip install -U psycopg2`."
- ) from e
- cursor = self.connection.cursor(cursor_factory=psycopg2.extras.NamedTupleCursor)
- cursor.execute("""LOAD 'age';""")
- cursor.execute("""SET search_path = ag_catalog, "$user", public;""")
- return cursor
-
- def _get_labels(self) -> Tuple[List[str], List[str]]:
- """
- Get all labels of a graph (for both edges and vertices)
- by querying the graph metadata table directly
-
- Returns
- Tuple[List[str]]: 2 lists, the first containing vertex
- labels and the second containing edge labels
- """
-
- e_labels_records = self.query(
- """MATCH ()-[e]-() RETURN collect(distinct label(e)) as labels"""
- )
- e_labels = e_labels_records[0]["labels"] if e_labels_records else []
-
- n_labels_records = self.query(
- """MATCH (n) RETURN collect(distinct label(n)) as labels"""
- )
- n_labels = n_labels_records[0]["labels"] if n_labels_records else []
-
- return n_labels, e_labels
-
- def _get_triples(self, e_labels: List[str]) -> List[Dict[str, str]]:
- """
- Get a set of distinct relationship types (as a list of dicts) in the graph
- to be used as context by an llm.
-
- Args:
- e_labels (List[str]): a list of edge labels to filter for
-
- Returns:
- List[Dict[str, str]]: relationships as a list of dicts in the format
- "{'start':, 'type':, 'end':}"
- """
-
- # age query to get distinct relationship types
- try:
- import psycopg2
- except ImportError as e:
- raise ImportError(
- "Unable to import psycopg2, please install with "
- "`pip install -U psycopg2`."
- ) from e
- triple_query = """
- SELECT * FROM ag_catalog.cypher('{graph_name}', $$
- MATCH (a)-[e:`{e_label}`]->(b)
- WITH a,e,b LIMIT 3000
- RETURN DISTINCT labels(a) AS from, type(e) AS edge, labels(b) AS to
- LIMIT 10
- $$) AS (f agtype, edge agtype, t agtype);
- """
-
- triple_schema = []
-
- # iterate desired edge types and add distinct relationship types to result
- with self._get_cursor() as curs:
- for label in e_labels:
- q = triple_query.format(graph_name=self.graph_name, e_label=label)
- try:
- curs.execute(q)
- data = curs.fetchall()
- for d in data:
- # use json.loads to convert returned
- # strings to python primitives
- triple_schema.append(
- {
- "start": json.loads(d.f)[0],
- "type": json.loads(d.edge),
- "end": json.loads(d.t)[0],
- }
- )
- except psycopg2.Error as e:
- raise AGEQueryException(
- {
- "message": "Error fetching triples",
- "detail": str(e),
- }
- )
-
- return triple_schema
-
- def _get_triples_str(self, e_labels: List[str]) -> List[str]:
- """
- Get a set of distinct relationship types (as a list of strings) in the graph
- to be used as context by an llm.
-
- Args:
- e_labels (List[str]): a list of edge labels to filter for
-
- Returns:
- List[str]: relationships as a list of strings in the format
- "(:``)-[:``]->(:``)"
- """
-
- triples = self._get_triples(e_labels)
-
- return self._format_triples(triples)
-
- @staticmethod
- def _format_triples(triples: List[Dict[str, str]]) -> List[str]:
- """
- Convert a list of relationships from dictionaries to formatted strings
- to be better readable by an llm
-
- Args:
- triples (List[Dict[str,str]]): a list relationships in the form
- {'start':, 'type':, 'end':}
-
- Returns:
- List[str]: a list of relationships in the form
- "(:``)-[:``]->(:``)"
- """
- triple_template = "(:`{start}`)-[:`{type}`]->(:`{end}`)"
- triple_schema = [triple_template.format(**triple) for triple in triples]
-
- return triple_schema
-
- def _get_node_properties(self, n_labels: List[str]) -> List[Dict[str, Any]]:
- """
- Fetch a list of available node properties by node label to be used
- as context for an llm
-
- Args:
- n_labels (List[str]): a list of node labels to filter for
-
- Returns:
- List[Dict[str, Any]]: a list of node labels and
- their corresponding properties in the form
- "{
- 'labels': ,
- 'properties': [
- {
- 'property': ,
- 'type':
- },...
- ]
- }"
- """
- try:
- import psycopg2
- except ImportError as e:
- raise ImportError(
- "Unable to import psycopg2, please install with "
- "`pip install -U psycopg2`."
- ) from e
-
- # cypher query to fetch properties of a given label
- node_properties_query = """
- SELECT * FROM ag_catalog.cypher('{graph_name}', $$
- MATCH (a:`{n_label}`)
- RETURN properties(a) AS props
- LIMIT 100
- $$) AS (props agtype);
- """
-
- node_properties = []
- with self._get_cursor() as curs:
- for label in n_labels:
- q = node_properties_query.format(
- graph_name=self.graph_name, n_label=label
- )
-
- try:
- curs.execute(q)
- except psycopg2.Error as e:
- raise AGEQueryException(
- {
- "message": "Error fetching node properties",
- "detail": str(e),
- }
- )
- data = curs.fetchall()
-
- # build a set of distinct properties
- s = set({})
- for d in data:
- # use json.loads to convert to python
- # primitive and get readable type
- for k, v in json.loads(d.props).items():
- s.add((k, self.types[type(v).__name__]))
-
- np = {
- "properties": [{"property": k, "type": v} for k, v in s],
- "labels": label,
- }
- node_properties.append(np)
-
- return node_properties
-
- def _get_edge_properties(self, e_labels: List[str]) -> List[Dict[str, Any]]:
- """
- Fetch a list of available edge properties by edge label to be used
- as context for an llm
-
- Args:
- e_labels (List[str]): a list of edge labels to filter for
-
- Returns:
- List[Dict[str, Any]]: a list of edge labels
- and their corresponding properties in the form
- "{
- 'labels': ,
- 'properties': [
- {
- 'property': ,
- 'type':
- },...
- ]
- }"
- """
-
- try:
- import psycopg2
- except ImportError as e:
- raise ImportError(
- "Unable to import psycopg2, please install with "
- "`pip install -U psycopg2`."
- ) from e
- # cypher query to fetch properties of a given label
- edge_properties_query = """
- SELECT * FROM ag_catalog.cypher('{graph_name}', $$
- MATCH ()-[e:`{e_label}`]->()
- RETURN properties(e) AS props
- LIMIT 100
- $$) AS (props agtype);
- """
- edge_properties = []
- with self._get_cursor() as curs:
- for label in e_labels:
- q = edge_properties_query.format(
- graph_name=self.graph_name, e_label=label
- )
-
- try:
- curs.execute(q)
- except psycopg2.Error as e:
- raise AGEQueryException(
- {
- "message": "Error fetching edge properties",
- "detail": str(e),
- }
- )
- data = curs.fetchall()
-
- # build a set of distinct properties
- s = set({})
- for d in data:
- # use json.loads to convert to python
- # primitive and get readable type
- for k, v in json.loads(d.props).items():
- s.add((k, self.types[type(v).__name__]))
-
- np = {
- "properties": [{"property": k, "type": v} for k, v in s],
- "type": label,
- }
- edge_properties.append(np)
-
- return edge_properties
-
- def refresh_schema(self) -> None:
- """
- Refresh the graph schema information by updating the available
- labels, relationships, and properties
- """
-
- # fetch graph schema information
- n_labels, e_labels = self._get_labels()
- triple_schema = self._get_triples(e_labels)
-
- node_properties = self._get_node_properties(n_labels)
- edge_properties = self._get_edge_properties(e_labels)
-
- # update the formatted string representation
- self.schema = f"""
- Node properties are the following:
- {node_properties}
- Relationship properties are the following:
- {edge_properties}
- The relationships are the following:
- {self._format_triples(triple_schema)}
- """
-
- # update the dictionary representation
- self.structured_schema = {
- "node_props": {el["labels"]: el["properties"] for el in node_properties},
- "rel_props": {el["type"]: el["properties"] for el in edge_properties},
- "relationships": triple_schema,
- "metadata": {},
- }
-
- @property
- def get_schema(self) -> str:
- """Returns the schema of the Graph"""
- return self.schema
-
- @property
- def get_structured_schema(self) -> Dict[str, Any]:
- """Returns the structured schema of the Graph"""
- return self.structured_schema
-
- @staticmethod
- def _get_col_name(field: str, idx: int) -> str:
- """
- Convert a cypher return field to a pgsql select field
- If possible keep the cypher column name, but create a generic name if necessary
-
- Args:
- field (str): a return field from a cypher query to be formatted for pgsql
- idx (int): the position of the field in the return statement
-
- Returns:
- str: the field to be used in the pgsql select statement
- """
- # remove white space
- field = field.strip()
- # if an alias is provided for the field, use it
- if " as " in field:
- return field.split(" as ")[-1].strip()
- # if the return value is an unnamed primitive, give it a generic name
- elif field.isnumeric() or field in ("true", "false", "null"):
- return f"column_{idx}"
- # otherwise return the value stripping out some common special chars
- else:
- return field.replace("(", "_").replace(")", "")
-
- @staticmethod
- def _wrap_query(query: str, graph_name: str) -> str:
- """
- Convert a Cyper query to an Apache Age compatible Sql Query.
- Handles combined queries with UNION/EXCEPT operators
-
- Args:
- query (str) : A valid cypher query, can include UNION/EXCEPT operators
- graph_name (str) : The name of the graph to query
-
- Returns :
- str : An equivalent pgSql query wrapped with ag_catalog.cypher
-
- Raises:
- ValueError : If query is empty, contain RETURN *, or has invalid field names
- """
-
- if not query.strip():
- raise ValueError("Empty query provided")
-
- # pgsql template
- template = """SELECT {projection} FROM ag_catalog.cypher('{graph_name}', $$
- {query}
- $$) AS ({fields});"""
-
- # split the query into parts based on UNION and EXCEPT
- parts = re.split(r"\b(UNION\b|\bEXCEPT)\b", query, flags=re.IGNORECASE)
-
- all_fields = []
-
- for part in parts:
- if part.strip().upper() in ("UNION", "EXCEPT"):
- continue
-
- # if there are any returned fields they must be added to the pgsql query
- return_match = re.search(r'\breturn\b(?![^"]*")', part, re.IGNORECASE)
- if return_match:
- # Extract the part of the query after the RETURN keyword
- return_clause = part[return_match.end() :]
-
- # parse return statement to identify returned fields
- fields = (
- return_clause.lower()
- .split("distinct")[-1]
- .split("order by")[0]
- .split("skip")[0]
- .split("limit")[0]
- .split(",")
- )
-
- # raise exception if RETURN * is found as we can't resolve the fields
- clean_fileds = [f.strip() for f in fields if f.strip()]
- if "*" in clean_fileds:
- raise ValueError(
- "Apache Age does not support RETURN * in Cypher queries"
- )
-
- # Format fields and maintain order of appearance
- for idx, field in enumerate(clean_fileds):
- field_name = AGEGraph._get_col_name(field, idx)
- if field_name not in all_fields:
- all_fields.append(field_name)
-
- # if no return statements found in any part
- if not all_fields:
- fields_str = "a agtype"
-
- else:
- fields_str = ", ".join(f"{field} agtype" for field in all_fields)
-
- return template.format(
- graph_name=graph_name,
- query=query,
- fields=fields_str,
- projection="*",
- )
-
- @staticmethod
- def _record_to_dict(record: NamedTuple) -> Dict[str, Any]:
- """
- Convert a record returned from an age query to a dictionary
-
- Args:
- record (): a record from an age query result
-
- Returns:
- Dict[str, Any]: a dictionary representation of the record where
- the dictionary key is the field name and the value is the
- value converted to a python type
- """
- # result holder
- d = {}
-
- # prebuild a mapping of vertex_id to vertex mappings to be used
- # later to build edges
- vertices = {}
- for k in record._fields:
- v = getattr(record, k)
- # agtype comes back '{key: value}::type' which must be parsed
- if isinstance(v, str) and "::" in v:
- dtype = v.split("::")[-1]
- v = v.split("::")[0]
- if dtype == "vertex":
- vertex = json.loads(v)
- vertices[vertex["id"]] = vertex.get("properties")
-
- # iterate returned fields and parse appropriately
- for k in record._fields:
- v = getattr(record, k)
- if isinstance(v, str) and "::" in v:
- dtype = v.split("::")[-1]
- v = v.split("::")[0]
- else:
- dtype = ""
-
- if dtype == "vertex":
- d[k] = json.loads(v).get("properties")
- # convert edge from id-label->id by replacing id with node information
- # we only do this if the vertex was also returned in the query
- # this is an attempt to be consistent with neo4j implementation
- elif dtype == "edge":
- edge = json.loads(v)
- d[k] = (
- vertices.get(edge["start_id"], {}),
- edge["label"],
- vertices.get(edge["end_id"], {}),
- )
- else:
- d[k] = json.loads(v) if isinstance(v, str) else v
-
- return d
-
- def query(self, query: str, params: dict = {}) -> List[Dict[str, Any]]:
- """
- Query the graph by taking a cypher query, converting it to an
- age compatible query, executing it and converting the result
-
- Args:
- query (str): a cypher query to be executed
- params (dict): parameters for the query (not used in this implementation)
-
- Returns:
- List[Dict[str, Any]]: a list of dictionaries containing the result set
- """
- try:
- import psycopg2
- except ImportError as e:
- raise ImportError(
- "Unable to import psycopg2, please install with "
- "`pip install -U psycopg2`."
- ) from e
-
- # convert cypher query to pgsql/age query
- wrapped_query = self._wrap_query(query, self.graph_name)
-
- # execute the query, rolling back on an error
- with self._get_cursor() as curs:
- try:
- curs.execute(wrapped_query)
- self.connection.commit()
- except psycopg2.Error as e:
- self.connection.rollback()
- raise AGEQueryException(
- {
- "message": "Error executing graph query: {}".format(query),
- "detail": str(e),
- }
- )
-
- data = curs.fetchall()
- if data is None:
- result = []
- # convert to dictionaries
- else:
- result = [self._record_to_dict(d) for d in data]
-
- return result
-
- @staticmethod
- def _format_properties(
- properties: Dict[str, Any], id: Union[str, None] = None
- ) -> str:
- """
- Convert a dictionary of properties to a string representation that
- can be used in a cypher query insert/merge statement.
-
- Args:
- properties (Dict[str,str]): a dictionary containing node/edge properties
- id (Union[str, None]): the id of the node or None if none exists
-
- Returns:
- str: the properties dictionary as a properly formatted string
- """
- props = []
- # wrap property key in backticks to escape
- for k, v in properties.items():
- prop = f"`{k}`: {json.dumps(v)}"
- props.append(prop)
- if id is not None and "id" not in properties:
- props.append(
- f"id: {json.dumps(id)}" if isinstance(id, str) else f"id: {id}"
- )
- return "{" + ", ".join(props) + "}"
-
- @staticmethod
- def clean_graph_labels(label: str) -> str:
- """
- remove any disallowed characters from a label and replace with '_'
-
- Args:
- label (str): the original label
-
- Returns:
- str: the sanitized version of the label
- """
- return re.sub(AGEGraph.label_regex, "_", label)
-
- def add_graph_documents(
- self, graph_documents: List[GraphDocument], include_source: bool = False
- ) -> None:
- """
- insert a list of graph documents into the graph
-
- Args:
- graph_documents (List[GraphDocument]): the list of documents to be inserted
- include_source (bool): if True add nodes for the sources
- with MENTIONS edges to the entities they mention
-
- Returns:
- None
- """
- # query for inserting nodes
- node_insert_query = (
- """
- MERGE (n:`{label}` {{`id`: "{id}"}})
- SET n = {properties}
- """
- if not include_source
- else """
- MERGE (n:`{label}` {properties})
- MERGE (d:Document {d_properties})
- MERGE (d)-[:MENTIONS]->(n)
- """
- )
-
- # query for inserting edges
- edge_insert_query = """
- MERGE (from:`{f_label}` {f_properties})
- MERGE (to:`{t_label}` {t_properties})
- MERGE (from)-[:`{r_label}` {r_properties}]->(to)
- """
- # iterate docs and insert them
- for doc in graph_documents:
- # if we are adding sources, create an id for the source
- if include_source:
- if not doc.source.metadata.get("id"):
- doc.source.metadata["id"] = md5(
- doc.source.page_content.encode("utf-8")
- ).hexdigest()
-
- # insert entity nodes
- for node in doc.nodes:
- node.properties["id"] = node.id
- if include_source:
- query = node_insert_query.format(
- label=node.type,
- properties=self._format_properties(node.properties),
- d_properties=self._format_properties(doc.source.metadata),
- )
- else:
- query = node_insert_query.format(
- label=AGEGraph.clean_graph_labels(node.type),
- properties=self._format_properties(node.properties),
- id=node.id,
- )
-
- self.query(query)
-
- # insert relationships
- for edge in doc.relationships:
- edge.source.properties["id"] = edge.source.id
- edge.target.properties["id"] = edge.target.id
- inputs = {
- "f_label": AGEGraph.clean_graph_labels(edge.source.type),
- "f_properties": self._format_properties(edge.source.properties),
- "t_label": AGEGraph.clean_graph_labels(edge.target.type),
- "t_properties": self._format_properties(edge.target.properties),
- "r_label": AGEGraph.clean_graph_labels(edge.type).upper(),
- "r_properties": self._format_properties(edge.properties),
- }
-
- query = edge_insert_query.format(**inputs)
- self.query(query)
diff --git a/libs/community/langchain_community/graphs/arangodb_graph.py b/libs/community/langchain_community/graphs/arangodb_graph.py
deleted file mode 100644
index 5e354e27d4..0000000000
--- a/libs/community/langchain_community/graphs/arangodb_graph.py
+++ /dev/null
@@ -1,182 +0,0 @@
-import os
-from math import ceil
-from typing import Any, Dict, List, Optional
-
-
-class ArangoGraph:
- """ArangoDB wrapper for graph operations.
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- def __init__(self, db: Any) -> None:
- """Create a new ArangoDB graph wrapper instance."""
- self.set_db(db)
- self.set_schema()
-
- @property
- def db(self) -> Any:
- return self.__db
-
- @property
- def schema(self) -> Dict[str, Any]:
- return self.__schema
-
- def set_db(self, db: Any) -> None:
- from arango.database import Database
-
- if not isinstance(db, Database):
- msg = "**db** parameter must inherit from arango.database.Database"
- raise TypeError(msg)
-
- self.__db: Database = db
- self.set_schema()
-
- def set_schema(self, schema: Optional[Dict[str, Any]] = None) -> None:
- """
- Set the schema of the ArangoDB Database.
- Auto-generates Schema if **schema** is None.
- """
- self.__schema = self.generate_schema() if schema is None else schema
-
- def generate_schema(
- self, sample_ratio: float = 0
- ) -> Dict[str, List[Dict[str, Any]]]:
- """
- Generates the schema of the ArangoDB Database and returns it
- User can specify a **sample_ratio** (0 to 1) to determine the
- ratio of documents/edges used (in relation to the Collection size)
- to render each Collection Schema.
- """
- if not 0 <= sample_ratio <= 1:
- raise ValueError("**sample_ratio** value must be in between 0 to 1")
-
- # Stores the Edge Relationships between each ArangoDB Document Collection
- graph_schema: List[Dict[str, Any]] = [
- {"graph_name": g["name"], "edge_definitions": g["edge_definitions"]}
- for g in self.db.graphs()
- ]
-
- # Stores the schema of every ArangoDB Document/Edge collection
- collection_schema: List[Dict[str, Any]] = []
-
- for collection in self.db.collections():
- if collection["system"]:
- continue
-
- # Extract collection name, type, and size
- col_name: str = collection["name"]
- col_type: str = collection["type"]
- col_size: int = self.db.collection(col_name).count()
-
- # Skip collection if empty
- if col_size == 0:
- continue
-
- # Set number of ArangoDB documents/edges to retrieve
- limit_amount = ceil(sample_ratio * col_size) or 1
-
- aql = f"""
- FOR doc in `{col_name}`
- LIMIT {limit_amount}
- RETURN doc
- """
-
- doc: Dict[str, Any]
- properties: List[Dict[str, str]] = []
- for doc in self.__db.aql.execute(aql):
- for key, value in doc.items():
- properties.append({"name": key, "type": type(value).__name__})
-
- collection_schema.append(
- {
- "collection_name": col_name,
- "collection_type": col_type,
- f"{col_type}_properties": properties,
- f"example_{col_type}": doc,
- }
- )
-
- return {"Graph Schema": graph_schema, "Collection Schema": collection_schema}
-
- def query(
- self, query: str, top_k: Optional[int] = None, **kwargs: Any
- ) -> List[Dict[str, Any]]:
- """Query the ArangoDB database."""
- import itertools
-
- cursor = self.__db.aql.execute(query, **kwargs)
- return [doc for doc in itertools.islice(cursor, top_k)]
-
- @classmethod
- def from_db_credentials(
- cls,
- url: Optional[str] = None,
- dbname: Optional[str] = None,
- username: Optional[str] = None,
- password: Optional[str] = None,
- ) -> Any:
- """Convenience constructor that builds Arango DB from credentials.
-
- Args:
- url: Arango DB url. Can be passed in as named arg or set as environment
- var ``ARANGODB_URL``. Defaults to "http://localhost:8529".
- dbname: Arango DB name. Can be passed in as named arg or set as
- environment var ``ARANGODB_DBNAME``. Defaults to "_system".
- username: Can be passed in as named arg or set as environment var
- ``ARANGODB_USERNAME``. Defaults to "root".
- password: Can be passed ni as named arg or set as environment var
- ``ARANGODB_PASSWORD``. Defaults to "".
-
- Returns:
- An arango.database.StandardDatabase.
- """
- db = get_arangodb_client(
- url=url, dbname=dbname, username=username, password=password
- )
- return cls(db)
-
-
-def get_arangodb_client(
- url: Optional[str] = None,
- dbname: Optional[str] = None,
- username: Optional[str] = None,
- password: Optional[str] = None,
-) -> Any:
- """Get the Arango DB client from credentials.
-
- Args:
- url: Arango DB url. Can be passed in as named arg or set as environment
- var ``ARANGODB_URL``. Defaults to "http://localhost:8529".
- dbname: Arango DB name. Can be passed in as named arg or set as
- environment var ``ARANGODB_DBNAME``. Defaults to "_system".
- username: Can be passed in as named arg or set as environment var
- ``ARANGODB_USERNAME``. Defaults to "root".
- password: Can be passed ni as named arg or set as environment var
- ``ARANGODB_PASSWORD``. Defaults to "".
-
- Returns:
- An arango.database.StandardDatabase.
- """
- try:
- from arango import ArangoClient
- except ImportError as e:
- raise ImportError(
- "Unable to import arango, please install with `pip install python-arango`."
- ) from e
-
- _url: str = url or os.environ.get("ARANGODB_URL", "http://localhost:8529") # type: ignore[assignment]
- _dbname: str = dbname or os.environ.get("ARANGODB_DBNAME", "_system") # type: ignore[assignment]
- _username: str = username or os.environ.get("ARANGODB_USERNAME", "root") # type: ignore[assignment]
- _password: str = password or os.environ.get("ARANGODB_PASSWORD", "") # type: ignore[assignment]
-
- return ArangoClient(_url).db(_dbname, _username, _password, verify=True)
diff --git a/libs/community/langchain_community/graphs/falkordb_graph.py b/libs/community/langchain_community/graphs/falkordb_graph.py
deleted file mode 100644
index 56ce03c1f9..0000000000
--- a/libs/community/langchain_community/graphs/falkordb_graph.py
+++ /dev/null
@@ -1,201 +0,0 @@
-import warnings
-from typing import Any, Dict, List, Optional
-
-from langchain_core._api import deprecated
-
-from langchain_community.graphs.graph_document import GraphDocument
-from langchain_community.graphs.graph_store import GraphStore
-
-node_properties_query = """
-MATCH (n)
-WITH keys(n) as keys, labels(n) AS labels
-WITH CASE WHEN keys = [] THEN [NULL] ELSE keys END AS keys, labels
-UNWIND labels AS label
-UNWIND keys AS key
-WITH label, collect(DISTINCT key) AS keys
-RETURN {label:label, keys:keys} AS output
-"""
-
-rel_properties_query = """
-MATCH ()-[r]->()
-WITH keys(r) as keys, type(r) AS types
-WITH CASE WHEN keys = [] THEN [NULL] ELSE keys END AS keys, types
-UNWIND types AS type
-UNWIND keys AS key WITH type,
-collect(DISTINCT key) AS keys
-RETURN {types:type, keys:keys} AS output
-"""
-
-rel_query = """
-MATCH (n)-[r]->(m)
-UNWIND labels(n) as src_label
-UNWIND labels(m) as dst_label
-UNWIND type(r) as rel_type
-RETURN DISTINCT {start: src_label, type: rel_type, end: dst_label} AS output
-"""
-
-
-class FalkorDBGraph(GraphStore):
- """FalkorDB wrapper for graph operations.
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- def __init__(
- self,
- database: str,
- host: str = "localhost",
- port: int = 6379,
- username: Optional[str] = None,
- password: Optional[str] = None,
- ssl: bool = False,
- ) -> None:
- """Create a new FalkorDB graph wrapper instance."""
- try:
- self.__init_falkordb_connection(
- database, host, port, username, password, ssl
- )
-
- except ImportError:
- try:
- # Falls back to using the redis package just for backwards compatibility
- self.__init_redis_connection(
- database, host, port, username, password, ssl
- )
- except ImportError:
- raise ImportError(
- "Could not import falkordb python package. "
- "Please install it with `pip install falkordb`."
- )
-
- self.schema: str = ""
- self.structured_schema: Dict[str, Any] = {}
-
- try:
- self.refresh_schema()
- except Exception as e:
- raise ValueError(f"Could not refresh schema. Error: {e}")
-
- def __init_falkordb_connection(
- self,
- database: str,
- host: str = "localhost",
- port: int = 6379,
- username: Optional[str] = None,
- password: Optional[str] = None,
- ssl: bool = False,
- ) -> None:
- from falkordb import FalkorDB
-
- try:
- self._driver = FalkorDB(
- host=host, port=port, username=username, password=password, ssl=ssl
- )
- except Exception as e:
- raise ConnectionError(f"Failed to connect to FalkorDB: {e}")
-
- self._graph = self._driver.select_graph(database)
-
- @deprecated("0.0.31", alternative="__init_falkordb_connection")
- def __init_redis_connection(
- self,
- database: str,
- host: str = "localhost",
- port: int = 6379,
- username: Optional[str] = None,
- password: Optional[str] = None,
- ssl: bool = False,
- ) -> None:
- import redis
- from redis.commands.graph import Graph
-
- # show deprecation warning
- warnings.warn(
- "Using the redis package is deprecated. "
- "Please use the falkordb package instead, "
- "install it with `pip install falkordb`.",
- DeprecationWarning,
- )
-
- self._driver = redis.Redis(
- host=host, port=port, username=username, password=password, ssl=ssl
- )
-
- self._graph = Graph(self._driver, database)
-
- @property
- def get_schema(self) -> str:
- """Returns the schema of the FalkorDB database"""
- return self.schema
-
- @property
- def get_structured_schema(self) -> Dict[str, Any]:
- """Returns the structured schema of the Graph"""
- return self.structured_schema
-
- def refresh_schema(self) -> None:
- """Refreshes the schema of the FalkorDB database"""
- node_properties: List[Any] = self.query(node_properties_query)
- rel_properties: List[Any] = self.query(rel_properties_query)
- relationships: List[Any] = self.query(rel_query)
-
- self.structured_schema = {
- "node_props": {el[0]["label"]: el[0]["keys"] for el in node_properties},
- "rel_props": {el[0]["types"]: el[0]["keys"] for el in rel_properties},
- "relationships": [el[0] for el in relationships],
- }
-
- self.schema = (
- f"Node properties: {node_properties}\n"
- f"Relationships properties: {rel_properties}\n"
- f"Relationships: {relationships}\n"
- )
-
- def query(self, query: str, params: dict = {}) -> List[Dict[str, Any]]:
- """Query FalkorDB database."""
-
- try:
- data = self._graph.query(query, params)
- return data.result_set
- except Exception as e:
- raise ValueError(f"Generated Cypher Statement is not valid\n{e}")
-
- def add_graph_documents(
- self, graph_documents: List[GraphDocument], include_source: bool = False
- ) -> None:
- """
- Take GraphDocument as input as uses it to construct a graph.
- """
- for document in graph_documents:
- # Import nodes
- for node in document.nodes:
- self.query(
- (
- f"MERGE (n:{node.type} {{id:'{node.id}'}}) "
- "SET n += $properties "
- "RETURN distinct 'done' AS result"
- ),
- {"properties": node.properties},
- )
-
- # Import relationships
- for rel in document.relationships:
- self.query(
- (
- f"MATCH (a:{rel.source.type} {{id:'{rel.source.id}'}}), "
- f"(b:{rel.target.type} {{id:'{rel.target.id}'}}) "
- f"MERGE (a)-[r:{(rel.type.replace(' ', '_').upper())}]->(b) "
- "SET r += $properties "
- "RETURN distinct 'done' AS result"
- ),
- {"properties": rel.properties},
- )
diff --git a/libs/community/langchain_community/graphs/graph_document.py b/libs/community/langchain_community/graphs/graph_document.py
deleted file mode 100644
index ff82ca4b43..0000000000
--- a/libs/community/langchain_community/graphs/graph_document.py
+++ /dev/null
@@ -1,51 +0,0 @@
-from __future__ import annotations
-
-from typing import List, Union
-
-from langchain_core.documents import Document
-from langchain_core.load.serializable import Serializable
-from pydantic import Field
-
-
-class Node(Serializable):
- """Represents a node in a graph with associated properties.
-
- Attributes:
- id (Union[str, int]): A unique identifier for the node.
- type (str): The type or label of the node, default is "Node".
- properties (dict): Additional properties and metadata associated with the node.
- """
-
- id: Union[str, int]
- type: str = "Node"
- properties: dict = Field(default_factory=dict)
-
-
-class Relationship(Serializable):
- """Represents a directed relationship between two nodes in a graph.
-
- Attributes:
- source (Node): The source node of the relationship.
- target (Node): The target node of the relationship.
- type (str): The type of the relationship.
- properties (dict): Additional properties associated with the relationship.
- """
-
- source: Node
- target: Node
- type: str
- properties: dict = Field(default_factory=dict)
-
-
-class GraphDocument(Serializable):
- """Represents a graph document consisting of nodes and relationships.
-
- Attributes:
- nodes (List[Node]): A list of nodes in the graph.
- relationships (List[Relationship]): A list of relationships in the graph.
- source (Document): The document from which the graph information is derived.
- """
-
- nodes: List[Node]
- relationships: List[Relationship]
- source: Document
diff --git a/libs/community/langchain_community/graphs/graph_store.py b/libs/community/langchain_community/graphs/graph_store.py
deleted file mode 100644
index 73a07c7de5..0000000000
--- a/libs/community/langchain_community/graphs/graph_store.py
+++ /dev/null
@@ -1,37 +0,0 @@
-from abc import abstractmethod
-from typing import Any, Dict, List
-
-from langchain_community.graphs.graph_document import GraphDocument
-
-
-class GraphStore:
- """Abstract class for graph operations."""
-
- @property
- @abstractmethod
- def get_schema(self) -> str:
- """Return the schema of the Graph database"""
- pass
-
- @property
- @abstractmethod
- def get_structured_schema(self) -> Dict[str, Any]:
- """Return the schema of the Graph database"""
- pass
-
- @abstractmethod
- def query(self, query: str, params: dict = {}) -> List[Dict[str, Any]]:
- """Query the graph."""
- pass
-
- @abstractmethod
- def refresh_schema(self) -> None:
- """Refresh the graph schema information."""
- pass
-
- @abstractmethod
- def add_graph_documents(
- self, graph_documents: List[GraphDocument], include_source: bool = False
- ) -> None:
- """Take GraphDocument as input as uses it to construct a graph."""
- pass
diff --git a/libs/community/langchain_community/graphs/gremlin_graph.py b/libs/community/langchain_community/graphs/gremlin_graph.py
deleted file mode 100644
index 26fe58eb1b..0000000000
--- a/libs/community/langchain_community/graphs/gremlin_graph.py
+++ /dev/null
@@ -1,228 +0,0 @@
-import hashlib
-import sys
-from typing import Any, Dict, List, Optional, Union
-
-from langchain_core.utils import get_from_env
-
-from langchain_community.graphs.graph_document import GraphDocument, Node, Relationship
-from langchain_community.graphs.graph_store import GraphStore
-
-
-class GremlinGraph(GraphStore):
- """Gremlin wrapper for graph operations.
-
- Parameters:
- url (Optional[str]): The URL of the Gremlin database server or env GREMLIN_URI
- username (Optional[str]): The collection-identifier like '/dbs/database/colls/graph'
- or env GREMLIN_USERNAME if none provided
- password (Optional[str]): The connection-key for database authentication
- or env GREMLIN_PASSWORD if none provided
- traversal_source (str): The traversal source to use for queries. Defaults to 'g'.
- message_serializer (Optional[Any]): The message serializer to use for requests.
- Defaults to serializer.GraphSONSerializersV2d0()
- include_edge_properties (bool): Whether to include edge properties in
- the gremlin graph schema. Defaults to False
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
-
- *Implementation details*:
- The Gremlin queries are designed to work with Azure CosmosDB limitations
- """
-
- @property
- def get_structured_schema(self) -> Dict[str, Any]:
- return self.structured_schema
-
- def __init__(
- self,
- url: Optional[str] = None,
- username: Optional[str] = None,
- password: Optional[str] = None,
- traversal_source: str = "g",
- message_serializer: Optional[Any] = None,
- include_edge_properties: bool = False,
- ) -> None:
- """Create a new Gremlin graph wrapper instance."""
- try:
- import asyncio
-
- from gremlin_python.driver import client, serializer
-
- if sys.platform == "win32":
- asyncio.set_event_loop_policy(asyncio.WindowsSelectorEventLoopPolicy())
- except ImportError:
- raise ImportError(
- "Please install gremlin-python first: `pip3 install gremlinpython"
- )
-
- self.client = client.Client(
- url=get_from_env("url", "GREMLIN_URI", url),
- traversal_source=traversal_source,
- username=get_from_env("username", "GREMLIN_USERNAME", username),
- password=get_from_env("password", "GREMLIN_PASSWORD", password),
- message_serializer=message_serializer
- if message_serializer
- else serializer.GraphSONSerializersV2d0(),
- )
- self.schema: str = ""
- self.include_edge_properties = include_edge_properties
-
- @property
- def get_schema(self) -> str:
- """Returns the schema of the Gremlin database"""
- if len(self.schema) == 0:
- self.refresh_schema()
- return self.schema
-
- def refresh_schema(self) -> None:
- """
- Refreshes the Gremlin graph schema information.
- """
- vertex_schema = self.client.submit("g.V().label().dedup()").all().result()
- edge_schema = self.client.submit("g.E().label().dedup()").all().result()
- vertex_properties = (
- self.client.submit(
- "g.V().group().by(label).by(properties().label().dedup().fold())"
- )
- .all()
- .result()[0]
- )
-
- self.structured_schema = {
- "vertex_labels": vertex_schema,
- "edge_labels": edge_schema,
- "vertice_props": vertex_properties,
- }
-
- self.schema = "\n".join(
- [
- "Vertex labels are the following:",
- ",".join(vertex_schema),
- "Edge labels are the following:",
- ",".join(edge_schema),
- f"Vertices have following properties:\n{vertex_properties}",
- ]
- )
- if self.include_edge_properties:
- edge_properties = (
- self.client.submit(
- "g.E().group().by(label)"
- ".by(project('inVLabel', 'outVLabel','properties')"
- ".by(inV().label()).by(outV().label()).by(properties().key().dedup()"
- ".fold()).dedup().fold())"
- )
- .all()
- .result()[0]
- )
- self.structured_schema["edge_props"] = edge_properties
- self.schema += (
- f"\nEdges have the following properties, grouped by label and"
- f" the distinct inV and outV labels:\n {edge_properties}"
- )
-
- def query(self, query: str, params: dict = {}) -> List[Dict[str, Any]]:
- q = self.client.submit(query)
- return q.all().result()
-
- def add_graph_documents(
- self, graph_documents: List[GraphDocument], include_source: bool = False
- ) -> None:
- """
- Take GraphDocument as input as uses it to construct a graph.
- """
- node_cache: Dict[Union[str, int], Node] = {}
- for document in graph_documents:
- if include_source:
- # Create document vertex
- doc_props = {
- "page_content": document.source.page_content,
- "metadata": document.source.metadata,
- }
- doc_id = hashlib.md5(document.source.page_content.encode()).hexdigest()
- doc_node = self.add_node(
- Node(id=doc_id, type="Document", properties=doc_props), node_cache
- )
-
- # Import nodes to vertices
- for n in document.nodes:
- node = self.add_node(n)
- if include_source:
- # Add Edge to document for each node
- self.add_edge(
- Relationship(
- type="contains information about",
- source=doc_node,
- target=node,
- properties={},
- )
- )
- self.add_edge(
- Relationship(
- type="is extracted from",
- source=node,
- target=doc_node,
- properties={},
- )
- )
-
- # Edges
- for el in document.relationships:
- # Find or create the source vertex
- self.add_node(el.source, node_cache)
- # Find or create the target vertex
- self.add_node(el.target, node_cache)
- # Find or create the edge
- self.add_edge(el)
-
- def build_vertex_query(self, node: Node) -> str:
- base_query = (
- f"g.V().has('id','{node.id}').fold()"
- + f".coalesce(unfold(),addV('{node.type}')"
- + f".property('id','{node.id}')"
- + f".property('type','{node.type}')"
- )
- for key, value in node.properties.items():
- base_query += f".property('{key}', '{value}')"
-
- return base_query + ")"
-
- def build_edge_query(self, relationship: Relationship) -> str:
- source_query = f".has('id','{relationship.source.id}')"
- target_query = f".has('id','{relationship.target.id}')"
-
- base_query = f""""g.V(){source_query}.as('a')
- .V(){target_query}.as('b')
- .choose(
- __.inE('{relationship.type}').where(outV().as('a')),
- __.identity(),
- __.addE('{relationship.type}').from('a').to('b')
- )
- """.replace("\n", "").replace("\t", "")
- for key, value in relationship.properties.items():
- base_query += f".property('{key}', '{value}')"
-
- return base_query
-
- def add_node(self, node: Node, node_cache: dict = {}) -> Node:
- # if properties does not have label, add type as label
- if "label" not in node.properties:
- node.properties["label"] = node.type
- if node.id in node_cache:
- return node_cache[node.id]
- else:
- query = self.build_vertex_query(node)
- _ = self.client.submit(query).all().result()[0]
- node_cache[node.id] = node
- return node
-
- def add_edge(self, relationship: Relationship) -> Any:
- query = self.build_edge_query(relationship)
- return self.client.submit(query).all().result()
diff --git a/libs/community/langchain_community/graphs/hugegraph.py b/libs/community/langchain_community/graphs/hugegraph.py
deleted file mode 100644
index 5bb6b167b0..0000000000
--- a/libs/community/langchain_community/graphs/hugegraph.py
+++ /dev/null
@@ -1,74 +0,0 @@
-from typing import Any, Dict, List
-
-
-class HugeGraph:
- """HugeGraph wrapper for graph operations.
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- def __init__(
- self,
- username: str = "default",
- password: str = "default",
- address: str = "127.0.0.1",
- port: int = 8081,
- graph: str = "hugegraph",
- ) -> None:
- """Create a new HugeGraph wrapper instance."""
- try:
- from hugegraph.connection import PyHugeGraph
- except ImportError:
- raise ImportError(
- "Please install HugeGraph Python client first: "
- "`pip3 install hugegraph-python`"
- )
-
- self.username = username
- self.password = password
- self.address = address
- self.port = port
- self.graph = graph
- self.client = PyHugeGraph(
- address, port, user=username, pwd=password, graph=graph
- )
- self.schema = ""
- # Set schema
- try:
- self.refresh_schema()
- except Exception as e:
- raise ValueError(f"Could not refresh schema. Error: {e}")
-
- @property
- def get_schema(self) -> str:
- """Returns the schema of the HugeGraph database"""
- return self.schema
-
- def refresh_schema(self) -> None:
- """
- Refreshes the HugeGraph schema information.
- """
- schema = self.client.schema()
- vertex_schema = schema.getVertexLabels()
- edge_schema = schema.getEdgeLabels()
- relationships = schema.getRelations()
-
- self.schema = (
- f"Node properties: {vertex_schema}\n"
- f"Edge properties: {edge_schema}\n"
- f"Relationships: {relationships}\n"
- )
-
- def query(self, query: str) -> List[Dict[str, Any]]:
- g = self.client.gremlin()
- res = g.exec(query)
- return res["data"]
diff --git a/libs/community/langchain_community/graphs/index_creator.py b/libs/community/langchain_community/graphs/index_creator.py
deleted file mode 100644
index f686096092..0000000000
--- a/libs/community/langchain_community/graphs/index_creator.py
+++ /dev/null
@@ -1,99 +0,0 @@
-from typing import Optional, Type
-
-
-from pydantic import BaseModel
-from langchain_core.language_models import BaseLanguageModel
-from langchain_core.prompts import BasePromptTemplate
-from langchain_core.prompts.prompt import PromptTemplate
-
-from langchain_community.graphs import NetworkxEntityGraph
-from langchain_community.graphs.networkx_graph import KG_TRIPLE_DELIMITER
-from langchain_community.graphs.networkx_graph import parse_triples
-
-# flake8: noqa
-
-_DEFAULT_KNOWLEDGE_TRIPLE_EXTRACTION_TEMPLATE = (
- "You are a networked intelligence helping a human track knowledge triples"
- " about all relevant people, things, concepts, etc. and integrating"
- " them with your knowledge stored within your weights"
- " as well as that stored in a knowledge graph."
- " Extract all of the knowledge triples from the text."
- " A knowledge triple is a clause that contains a subject, a predicate,"
- " and an object. The subject is the entity being described,"
- " the predicate is the property of the subject that is being"
- " described, and the object is the value of the property.\n\n"
- "EXAMPLE\n"
- "It's a state in the US. It's also the number 1 producer of gold in the US.\n\n"
- f"Output: (Nevada, is a, state){KG_TRIPLE_DELIMITER}(Nevada, is in, US)"
- f"{KG_TRIPLE_DELIMITER}(Nevada, is the number 1 producer of, gold)\n"
- "END OF EXAMPLE\n\n"
- "EXAMPLE\n"
- "I'm going to the store.\n\n"
- "Output: NONE\n"
- "END OF EXAMPLE\n\n"
- "EXAMPLE\n"
- "Oh huh. I know Descartes likes to drive antique scooters and play the mandolin.\n"
- f"Output: (Descartes, likes to drive, antique scooters){KG_TRIPLE_DELIMITER}(Descartes, plays, mandolin)\n"
- "END OF EXAMPLE\n\n"
- "EXAMPLE\n"
- "{text}"
- "Output:"
-)
-
-KNOWLEDGE_TRIPLE_EXTRACTION_PROMPT = PromptTemplate(
- input_variables=["text"],
- template=_DEFAULT_KNOWLEDGE_TRIPLE_EXTRACTION_TEMPLATE,
-)
-
-
-class GraphIndexCreator(BaseModel):
- """Functionality to create graph index."""
-
- llm: Optional[BaseLanguageModel] = None
- graph_type: Type[NetworkxEntityGraph] = NetworkxEntityGraph
-
- def from_text(
- self, text: str, prompt: BasePromptTemplate = KNOWLEDGE_TRIPLE_EXTRACTION_PROMPT
- ) -> NetworkxEntityGraph:
- """Create graph index from text."""
- if self.llm is None:
- raise ValueError("llm should not be None")
- graph = self.graph_type()
- # Temporary local scoped import while community does not depend on
- # langchain explicitly
- try:
- from langchain.chains import LLMChain
- except ImportError:
- raise ImportError(
- "Please install langchain to use this functionality. "
- "You can install it with `pip install langchain`."
- )
- chain = LLMChain(llm=self.llm, prompt=prompt)
- output = chain.predict(text=text)
- knowledge = parse_triples(output)
- for triple in knowledge:
- graph.add_triple(triple)
- return graph
-
- async def afrom_text(
- self, text: str, prompt: BasePromptTemplate = KNOWLEDGE_TRIPLE_EXTRACTION_PROMPT
- ) -> NetworkxEntityGraph:
- """Create graph index from text asynchronously."""
- if self.llm is None:
- raise ValueError("llm should not be None")
- graph = self.graph_type()
- # Temporary local scoped import while community does not depend on
- # langchain explicitly
- try:
- from langchain.chains import LLMChain
- except ImportError:
- raise ImportError(
- "Please install langchain to use this functionality. "
- "You can install it with `pip install langchain`."
- )
- chain = LLMChain(llm=self.llm, prompt=prompt)
- output = await chain.apredict(text=text)
- knowledge = parse_triples(output)
- for triple in knowledge:
- graph.add_triple(triple)
- return graph
diff --git a/libs/community/langchain_community/graphs/kuzu_graph.py b/libs/community/langchain_community/graphs/kuzu_graph.py
deleted file mode 100644
index b658d9510d..0000000000
--- a/libs/community/langchain_community/graphs/kuzu_graph.py
+++ /dev/null
@@ -1,264 +0,0 @@
-from hashlib import md5
-from typing import Any, Dict, List, Tuple
-
-from langchain_community.graphs.graph_document import GraphDocument, Relationship
-
-
-class KuzuGraph:
- """Kùzu wrapper for graph operations.
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- def __init__(
- self, db: Any, database: str = "kuzu", allow_dangerous_requests: bool = False
- ) -> None:
- """Initializes the Kùzu graph database connection."""
-
- if allow_dangerous_requests is not True:
- raise ValueError(
- "The KuzuGraph class is a powerful tool that can be used to execute "
- "arbitrary queries on the database. To enable this functionality, "
- "set the `allow_dangerous_requests` parameter to `True` when "
- "constructing the KuzuGraph object."
- )
-
- try:
- import kuzu
- except ImportError:
- raise ImportError(
- "Could not import Kùzu python package."
- "Please install Kùzu with `pip install kuzu`."
- )
- self.db = db
- self.conn = kuzu.Connection(self.db)
- self.database = database
- self.refresh_schema()
-
- @property
- def get_schema(self) -> str:
- """Returns the schema of the Kùzu database"""
- return self.schema
-
- def query(self, query: str, params: dict = {}) -> List[Dict[str, Any]]:
- """Query Kùzu database"""
- result = self.conn.execute(query, params)
- column_names = result.get_column_names()
- return_list = []
- while result.has_next():
- row = result.get_next()
- return_list.append(dict(zip(column_names, row)))
- return return_list
-
- def refresh_schema(self) -> None:
- """Refreshes the Kùzu graph schema information"""
- node_properties = []
- node_table_names = self.conn._get_node_table_names()
- for table_name in node_table_names:
- current_table_schema = {"properties": [], "label": table_name}
- properties = self.conn._get_node_property_names(table_name)
- for property_name in properties:
- property_type = properties[property_name]["type"]
- list_type_flag = ""
- if properties[property_name]["dimension"] > 0:
- if "shape" in properties[property_name]:
- for s in properties[property_name]["shape"]:
- list_type_flag += f"[{s}]"
- else:
- for i in range(properties[property_name]["dimension"]):
- list_type_flag += "[]"
- property_type += list_type_flag
- current_table_schema["properties"].append(
- (
- property_name,
- property_type,
- )
- )
- node_properties.append(current_table_schema)
-
- relationships = []
- rel_tables = self.conn._get_rel_table_names()
- for table in rel_tables:
- relationships.append(
- f"(:{table['src']})-[:{table['name']}]->(:{table['dst']})"
- )
-
- rel_properties = []
- for table in rel_tables:
- table_name = table["name"]
- current_table_schema = {"properties": [], "label": table_name}
- query_result = self.conn.execute(
- f"CALL table_info('{table_name}') RETURN *;"
- )
- while query_result.has_next():
- row = query_result.get_next()
- prop_name = row[1]
- prop_type = row[2]
- current_table_schema["properties"].append((prop_name, prop_type))
- rel_properties.append(current_table_schema)
-
- self.schema = (
- f"Node properties: {node_properties}\n"
- f"Relationships properties: {rel_properties}\n"
- f"Relationships: {relationships}\n"
- )
-
- def _create_chunk_node_table(self) -> None:
- self.conn.execute(
- """
- CREATE NODE TABLE IF NOT EXISTS Chunk (
- id STRING,
- text STRING,
- type STRING,
- PRIMARY KEY(id)
- );
- """
- )
-
- def _create_entity_node_table(self, node_label: str) -> None:
- self.conn.execute(
- f"""
- CREATE NODE TABLE IF NOT EXISTS {node_label} (
- id STRING,
- type STRING,
- PRIMARY KEY(id)
- );
- """
- )
-
- def _create_entity_relationship_table(self, rel: Relationship) -> None:
- self.conn.execute(
- f"""
- CREATE REL TABLE IF NOT EXISTS {rel.type} (
- FROM {rel.source.type} TO {rel.target.type}
- );
- """
- )
-
- def add_graph_documents(
- self,
- graph_documents: List[GraphDocument],
- allowed_relationships: List[Tuple[str, str, str]],
- include_source: bool = False,
- ) -> None:
- """
- Adds a list of `GraphDocument` objects that represent nodes and relationships
- in a graph to a Kùzu backend.
-
- Parameters:
- - graph_documents (List[GraphDocument]): A list of `GraphDocument` objects
- that contain the nodes and relationships to be added to the graph. Each
- `GraphDocument` should encapsulate the structure of part of the graph,
- including nodes, relationships, and the source document information.
-
- - allowed_relationships (List[Tuple[str, str, str]]): A list of allowed
- relationships that exist in the graph. Each tuple contains three elements:
- the source node type, the relationship type, and the target node type.
- Required for Kùzu, as the names of the relationship tables that need to
- pre-exist are derived from these tuples.
-
- - include_source (bool): If True, stores the source document
- and links it to nodes in the graph using the `MENTIONS` relationship.
- This is useful for tracing back the origin of data. Merges source
- documents based on the `id` property from the source document metadata
- if available; otherwise it calculates the MD5 hash of `page_content`
- for merging process. Defaults to False.
- """
- # Get unique node labels in the graph documents
- node_labels = list(
- {node.type for document in graph_documents for node in document.nodes}
- )
-
- for document in graph_documents:
- # Add chunk nodes and create source document relationships if include_source
- # is True
- if include_source:
- self._create_chunk_node_table()
- if not document.source.metadata.get("id"):
- # Add a unique id to each document chunk via an md5 hash
- document.source.metadata["id"] = md5(
- document.source.page_content.encode("utf-8")
- ).hexdigest()
-
- self.conn.execute(
- f"""
- MERGE (c:Chunk {{id: $id}})
- SET c.text = $text,
- c.type = "text_chunk"
- """, # noqa: F541
- parameters={
- "id": document.source.metadata["id"],
- "text": document.source.page_content,
- },
- )
-
- for node_label in node_labels:
- self._create_entity_node_table(node_label)
-
- # Add entity nodes from data
- for node in document.nodes:
- self.conn.execute(
- f"""
- MERGE (e:{node.type} {{id: $id}})
- SET e.type = "entity"
- """,
- parameters={"id": node.id},
- )
- if include_source:
- # If include_source is True, we need to create a relationship table
- # between the chunk nodes and the entity nodes
- self._create_chunk_node_table()
- ddl = "CREATE REL TABLE GROUP IF NOT EXISTS MENTIONS ("
- table_names = []
- for node_label in node_labels:
- table_names.append(f"FROM Chunk TO {node_label}")
- table_names = list(set(table_names))
- ddl += ", ".join(table_names)
- # Add common properties for all the tables here
- ddl += ", label STRING, triplet_source_id STRING)"
- if ddl:
- self.conn.execute(ddl)
-
- # Only allow relationships that exist in the schema
- if node.type in node_labels:
- self.conn.execute(
- f"""
- MATCH (c:Chunk {{id: $id}}),
- (e:{node.type} {{id: $node_id}})
- MERGE (c)-[m:MENTIONS]->(e)
- SET m.triplet_source_id = $id
- """,
- parameters={
- "id": document.source.metadata["id"],
- "node_id": node.id,
- },
- )
-
- # Add entity relationships
- for rel in document.relationships:
- self._create_entity_relationship_table(rel)
- # Create relationship
- source_label = rel.source.type
- source_id = rel.source.id
- target_label = rel.target.type
- target_id = rel.target.id
- self.conn.execute(
- f"""
- MATCH (e1:{source_label} {{id: $source_id}}),
- (e2:{target_label} {{id: $target_id}})
- MERGE (e1)-[:{rel.type}]->(e2)
- """,
- parameters={
- "source_id": source_id,
- "target_id": target_id,
- },
- )
diff --git a/libs/community/langchain_community/graphs/memgraph_graph.py b/libs/community/langchain_community/graphs/memgraph_graph.py
deleted file mode 100644
index 4180b49ce3..0000000000
--- a/libs/community/langchain_community/graphs/memgraph_graph.py
+++ /dev/null
@@ -1,525 +0,0 @@
-import logging
-from hashlib import md5
-from typing import Any, Dict, List, Optional
-
-from langchain_core.utils import get_from_dict_or_env
-
-from langchain_community.graphs.graph_document import GraphDocument, Node, Relationship
-from langchain_community.graphs.graph_store import GraphStore
-
-logger = logging.getLogger(__name__)
-
-
-BASE_ENTITY_LABEL = "__Entity__"
-
-SCHEMA_QUERY = """
-SHOW SCHEMA INFO
-"""
-
-NODE_PROPERTIES_QUERY = """
-CALL schema.node_type_properties()
-YIELD nodeType AS label, propertyName AS property, propertyTypes AS type
-WITH label AS nodeLabels, collect({key: property, types: type}) AS properties
-RETURN {labels: nodeLabels, properties: properties} AS output
-"""
-
-REL_QUERY = """
-MATCH (n)-[e]->(m)
-WITH DISTINCT
- labels(n) AS start_node_labels,
- type(e) AS rel_type,
- labels(m) AS end_node_labels,
- e,
- keys(e) AS properties
-UNWIND CASE WHEN size(properties) > 0 THEN properties ELSE [null] END AS prop
-WITH
- start_node_labels,
- rel_type,
- end_node_labels,
- CASE WHEN prop IS NULL THEN [] ELSE [prop, valueType(e[prop])] END AS property_info
-RETURN
- start_node_labels,
- rel_type,
- end_node_labels,
- COLLECT(DISTINCT CASE
- WHEN property_info <> []
- THEN property_info
- ELSE null END) AS properties_info
-"""
-
-NODE_IMPORT_QUERY = """
-UNWIND $data AS row
-CALL merge.node(row.label, row.properties, {}, {})
-YIELD node
-RETURN distinct 'done' AS result
-"""
-
-REL_NODES_IMPORT_QUERY = """
-UNWIND $data AS row
-MERGE (source {id: row.source_id})
-MERGE (target {id: row.target_id})
-RETURN distinct 'done' AS result
-"""
-
-REL_IMPORT_QUERY = """
-UNWIND $data AS row
-MATCH (source {id: row.source_id})
-MATCH (target {id: row.target_id})
-WITH source, target, row
-CALL merge.relationship(source, row.type, {}, {}, target, {})
-YIELD rel
-RETURN distinct 'done' AS result
-"""
-
-INCLUDE_DOCS_QUERY = """
-MERGE (d:Document {id:$document.metadata.id})
-SET d.content = $document.page_content
-SET d += $document.metadata
-RETURN distinct 'done' AS result
-"""
-
-INCLUDE_DOCS_SOURCE_QUERY = """
-UNWIND $data AS row
-MATCH (source {id: row.source_id}), (d:Document {id: $document.metadata.id})
-MERGE (d)-[:MENTIONS]->(source)
-RETURN distinct 'done' AS result
-"""
-
-NODE_PROPS_TEXT = """
-Node labels and properties (name and type) are:
-"""
-
-REL_PROPS_TEXT = """
-Relationship labels and properties are:
-"""
-
-REL_TEXT = """
-Nodes are connected with the following relationships:
-"""
-
-
-def get_schema_subset(data: Dict[str, Any]) -> Dict[str, Any]:
- return {
- "edges": [
- {
- "end_node_labels": edge["end_node_labels"],
- "properties": [
- {
- "key": prop["key"],
- "types": [
- {"type": type_item["type"].lower()}
- for type_item in prop["types"]
- ],
- }
- for prop in edge["properties"]
- ],
- "start_node_labels": edge["start_node_labels"],
- "type": edge["type"],
- }
- for edge in data["edges"]
- ],
- "nodes": [
- {
- "labels": node["labels"],
- "properties": [
- {
- "key": prop["key"],
- "types": [
- {"type": type_item["type"].lower()}
- for type_item in prop["types"]
- ],
- }
- for prop in node["properties"]
- ],
- }
- for node in data["nodes"]
- ],
- }
-
-
-def get_reformated_schema(
- nodes: List[Dict[str, Any]], rels: List[Dict[str, Any]]
-) -> Dict[str, Any]:
- return {
- "edges": [
- {
- "end_node_labels": rel["end_node_labels"],
- "properties": [
- {"key": prop[0], "types": [{"type": prop[1].lower()}]}
- for prop in rel["properties_info"]
- ],
- "start_node_labels": rel["start_node_labels"],
- "type": rel["rel_type"],
- }
- for rel in rels
- ],
- "nodes": [
- {
- "labels": [_remove_backticks(node["labels"])[1:]],
- "properties": [
- {
- "key": prop["key"],
- "types": [
- {"type": type_item.lower()} for type_item in prop["types"]
- ],
- }
- for prop in node["properties"]
- if node["properties"][0]["key"] != ""
- ],
- }
- for node in nodes
- ],
- }
-
-
-def transform_schema_to_text(schema: Dict[str, Any]) -> str:
- node_props_data = ""
- rel_props_data = ""
- rel_data = ""
-
- for node in schema["nodes"]:
- node_props_data += f"- labels: (:{':'.join(node['labels'])})\n"
- if node["properties"] == []:
- continue
- node_props_data += " properties:\n"
- for prop in node["properties"]:
- prop_types_str = " or ".join(
- {prop_types["type"] for prop_types in prop["types"]}
- )
- node_props_data += f" - {prop['key']}: {prop_types_str}\n"
-
- for rel in schema["edges"]:
- rel_type = rel["type"]
- start_labels = ":".join(rel["start_node_labels"])
- end_labels = ":".join(rel["end_node_labels"])
- rel_data += f"(:{start_labels})-[:{rel_type}]->(:{end_labels})\n"
-
- if rel["properties"] == []:
- continue
-
- rel_props_data += f"- labels: {rel_type}\n properties:\n"
- for prop in rel["properties"]:
- prop_types_str = " or ".join(
- {prop_types["type"].lower() for prop_types in prop["types"]}
- )
- rel_props_data += f" - {prop['key']}: {prop_types_str}\n"
-
- return "".join(
- [
- NODE_PROPS_TEXT + node_props_data if node_props_data else "",
- REL_PROPS_TEXT + rel_props_data if rel_props_data else "",
- REL_TEXT + rel_data if rel_data else "",
- ]
- )
-
-
-def _remove_backticks(text: str) -> str:
- return text.replace("`", "")
-
-
-def _transform_nodes(nodes: list[Node], baseEntityLabel: bool) -> List[dict]:
- transformed_nodes = []
- for node in nodes:
- properties_dict = node.properties | {"id": node.id}
- label = (
- [_remove_backticks(node.type), BASE_ENTITY_LABEL]
- if baseEntityLabel
- else [_remove_backticks(node.type)]
- )
- node_dict = {"label": label, "properties": properties_dict}
- transformed_nodes.append(node_dict)
- return transformed_nodes
-
-
-def _transform_relationships(
- relationships: list[Relationship], baseEntityLabel: bool
-) -> List[dict]:
- transformed_relationships = []
- for rel in relationships:
- rel_dict = {
- "type": _remove_backticks(rel.type),
- "source_label": (
- [BASE_ENTITY_LABEL]
- if baseEntityLabel
- else [_remove_backticks(rel.source.type)]
- ),
- "source_id": rel.source.id,
- "target_label": (
- [BASE_ENTITY_LABEL]
- if baseEntityLabel
- else [_remove_backticks(rel.target.type)]
- ),
- "target_id": rel.target.id,
- }
- transformed_relationships.append(rel_dict)
- return transformed_relationships
-
-
-class MemgraphGraph(GraphStore):
- """Memgraph wrapper for graph operations.
-
- Parameters:
- url (Optional[str]): The URL of the Memgraph database server.
- username (Optional[str]): The username for database authentication.
- password (Optional[str]): The password for database authentication.
- database (str): The name of the database to connect to. Default is 'memgraph'.
- refresh_schema (bool): A flag whether to refresh schema information
- at initialization. Default is True.
- driver_config (Dict): Configuration passed to Neo4j Driver.
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- def __init__(
- self,
- url: Optional[str] = None,
- username: Optional[str] = None,
- password: Optional[str] = None,
- database: Optional[str] = None,
- refresh_schema: bool = True,
- *,
- driver_config: Optional[Dict] = None,
- ) -> None:
- """Create a new Memgraph graph wrapper instance."""
- try:
- import neo4j
- except ImportError:
- raise ImportError(
- "Could not import neo4j python package. "
- "Please install it with `pip install neo4j`."
- )
-
- url = get_from_dict_or_env({"url": url}, "url", "MEMGRAPH_URI")
-
- # if username and password are "", assume auth is disabled
- if username == "" and password == "":
- auth = None
- else:
- username = get_from_dict_or_env(
- {"username": username},
- "username",
- "MEMGRAPH_USERNAME",
- )
- password = get_from_dict_or_env(
- {"password": password},
- "password",
- "MEMGRAPH_PASSWORD",
- )
- auth = (username, password)
- database = get_from_dict_or_env(
- {"database": database}, "database", "MEMGRAPH_DATABASE", "memgraph"
- )
-
- self._driver = neo4j.GraphDatabase.driver(
- url, auth=auth, **(driver_config or {})
- )
-
- self._database = database
- self.schema: str = ""
- self.structured_schema: Dict[str, Any] = {}
-
- # Verify connection
- try:
- self._driver.verify_connectivity()
- except neo4j.exceptions.ServiceUnavailable:
- raise ValueError(
- "Could not connect to Memgraph database. "
- "Please ensure that the url is correct"
- )
- except neo4j.exceptions.AuthError:
- raise ValueError(
- "Could not connect to Memgraph database. "
- "Please ensure that the username and password are correct"
- )
-
- # Set schema
- if refresh_schema:
- try:
- self.refresh_schema()
- except neo4j.exceptions.ClientError as e:
- raise e
-
- def close(self) -> None:
- if self._driver:
- logger.info("Closing the driver connection.")
- self._driver.close()
- self._driver = None
-
- @property
- def get_schema(self) -> str:
- """Returns the schema of the Graph database"""
- return self.schema
-
- @property
- def get_structured_schema(self) -> Dict[str, Any]:
- """Returns the structured schema of the Graph database"""
- return self.structured_schema
-
- def query(self, query: str, params: dict = {}) -> List[Dict[str, Any]]:
- """Query the graph.
-
- Args:
- query (str): The Cypher query to execute.
- params (dict): The parameters to pass to the query.
-
- Returns:
- List[Dict[str, Any]]: The list of dictionaries containing the query results.
- """
- from neo4j.exceptions import Neo4jError
-
- try:
- data, _, _ = self._driver.execute_query(
- query,
- database_=self._database,
- parameters_=params,
- )
- json_data = [r.data() for r in data]
- return json_data
- except Neo4jError as e:
- if not (
- (
- ( # isCallInTransactionError
- e.code == "Neo.DatabaseError.Statement.ExecutionFailed"
- or e.code
- == "Neo.DatabaseError.Transaction.TransactionStartFailed"
- )
- and "in an implicit transaction" in e.message
- )
- or ( # isPeriodicCommitError
- e.code == "Neo.ClientError.Statement.SemanticError"
- and (
- "in an open transaction is not possible" in e.message
- or "tried to execute in an explicit transaction" in e.message
- )
- )
- or (
- e.code == "Memgraph.ClientError.MemgraphError.MemgraphError"
- and ("in multicommand transactions" in e.message)
- )
- or (
- e.code == "Memgraph.ClientError.MemgraphError.MemgraphError"
- and "SchemaInfo disabled" in e.message
- )
- ):
- raise
-
- # fallback to allow implicit transactions
- with self._driver.session(database=self._database) as session:
- data = session.run(query, params)
- json_data = [r.data() for r in data]
- return json_data
-
- def refresh_schema(self) -> None:
- """
- Refreshes the Memgraph graph schema information.
- """
- import ast
-
- from neo4j.exceptions import Neo4jError
-
- # leave schema empty if db is empty
- if self.query("MATCH (n) RETURN n LIMIT 1") == []:
- return
-
- # first try with SHOW SCHEMA INFO
- try:
- result = self.query(SCHEMA_QUERY)[0].get("schema")
- if result is not None and isinstance(result, (str, ast.AST)):
- schema_result = ast.literal_eval(result)
- else:
- schema_result = result
- assert schema_result is not None
- structured_schema = get_schema_subset(schema_result)
- self.structured_schema = structured_schema
- self.schema = transform_schema_to_text(structured_schema)
- return
- except Neo4jError as e:
- if (
- e.code == "Memgraph.ClientError.MemgraphError.MemgraphError"
- and "SchemaInfo disabled" in e.message
- ):
- logger.info(
- "Schema generation with SHOW SCHEMA INFO query failed. "
- "Set --schema-info-enabled=true to use SHOW SCHEMA INFO query. "
- "Falling back to alternative queries."
- )
-
- # fallback on Cypher without SHOW SCHEMA INFO
- nodes = [query["output"] for query in self.query(NODE_PROPERTIES_QUERY)]
- rels = self.query(REL_QUERY)
-
- structured_schema = get_reformated_schema(nodes, rels)
- self.structured_schema = structured_schema
- self.schema = transform_schema_to_text(structured_schema)
-
- def add_graph_documents(
- self,
- graph_documents: List[GraphDocument],
- include_source: bool = False,
- baseEntityLabel: bool = False,
- ) -> None:
- """
- Take GraphDocument as input as uses it to construct a graph in Memgraph.
-
- Parameters:
- - graph_documents (List[GraphDocument]): A list of GraphDocument objects
- that contain the nodes and relationships to be added to the graph. Each
- GraphDocument should encapsulate the structure of part of the graph,
- including nodes, relationships, and the source document information.
- - include_source (bool, optional): If True, stores the source document
- and links it to nodes in the graph using the MENTIONS relationship.
- This is useful for tracing back the origin of data. Merges source
- documents based on the `id` property from the source document metadata
- if available; otherwise it calculates the MD5 hash of `page_content`
- for merging process. Defaults to False.
- - baseEntityLabel (bool, optional): If True, each newly created node
- gets a secondary __Entity__ label, which is indexed and improves import
- speed and performance. Defaults to False.
- """
-
- if baseEntityLabel:
- self.query(
- f"CREATE CONSTRAINT ON (b:{BASE_ENTITY_LABEL}) ASSERT b.id IS UNIQUE;"
- )
- self.query(f"CREATE INDEX ON :{BASE_ENTITY_LABEL}(id);")
- self.query(f"CREATE INDEX ON :{BASE_ENTITY_LABEL};")
-
- for document in graph_documents:
- if include_source:
- if not document.source.metadata.get("id"):
- document.source.metadata["id"] = md5(
- document.source.page_content.encode("utf-8")
- ).hexdigest()
-
- self.query(INCLUDE_DOCS_QUERY, {"document": document.source.__dict__})
-
- self.query(
- NODE_IMPORT_QUERY,
- {"data": _transform_nodes(document.nodes, baseEntityLabel)},
- )
-
- rel_data = _transform_relationships(document.relationships, baseEntityLabel)
- self.query(
- REL_NODES_IMPORT_QUERY,
- {"data": rel_data},
- )
- self.query(
- REL_IMPORT_QUERY,
- {"data": rel_data},
- )
-
- if include_source:
- self.query(
- INCLUDE_DOCS_SOURCE_QUERY,
- {"data": rel_data, "document": document.source.__dict__},
- )
- self.refresh_schema()
diff --git a/libs/community/langchain_community/graphs/nebula_graph.py b/libs/community/langchain_community/graphs/nebula_graph.py
deleted file mode 100644
index 81634fd6ea..0000000000
--- a/libs/community/langchain_community/graphs/nebula_graph.py
+++ /dev/null
@@ -1,222 +0,0 @@
-import logging
-from string import Template
-from typing import Any, Dict, Optional
-
-logger = logging.getLogger(__name__)
-
-rel_query = Template(
- """
-MATCH ()-[e:`$edge_type`]->()
- WITH e limit 1
-MATCH (m)-[:`$edge_type`]->(n) WHERE id(m) == src(e) AND id(n) == dst(e)
-RETURN "(:" + tags(m)[0] + ")-[:$edge_type]->(:" + tags(n)[0] + ")" AS rels
-"""
-)
-
-RETRY_TIMES = 3
-
-
-class NebulaGraph:
- """NebulaGraph wrapper for graph operations.
-
- NebulaGraph inherits methods from Neo4jGraph to bring ease to the user space.
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- def __init__(
- self,
- space: str,
- username: str = "root",
- password: str = "nebula",
- address: str = "127.0.0.1",
- port: int = 9669,
- session_pool_size: int = 30,
- ) -> None:
- """Create a new NebulaGraph wrapper instance."""
- try:
- import nebula3 # noqa: F401
- import pandas # noqa: F401
- except ImportError:
- raise ImportError(
- "Please install NebulaGraph Python client and pandas first: "
- "`pip install nebula3-python pandas`"
- )
-
- self.username = username
- self.password = password
- self.address = address
- self.port = port
- self.space = space
- self.session_pool_size = session_pool_size
-
- self.session_pool = self._get_session_pool()
- self.schema = ""
- # Set schema
- try:
- self.refresh_schema()
- except Exception as e:
- raise ValueError(f"Could not refresh schema. Error: {e}")
-
- def _get_session_pool(self) -> Any:
- assert all(
- [
- self.username,
- self.password,
- self.address,
- self.port,
- self.space,
- ]
- ), (
- "Please provide all of the following parameters: "
- "username, password, address, port, space"
- )
-
- from nebula3.Config import SessionPoolConfig
- from nebula3.Exception import AuthFailedException, InValidHostname
- from nebula3.gclient.net.SessionPool import SessionPool
-
- config = SessionPoolConfig()
- config.max_size = self.session_pool_size
-
- try:
- session_pool = SessionPool(
- self.username,
- self.password,
- self.space,
- [(self.address, self.port)],
- )
- except InValidHostname:
- raise ValueError(
- "Could not connect to NebulaGraph database. "
- "Please ensure that the address and port are correct"
- )
-
- try:
- session_pool.init(config)
- except AuthFailedException:
- raise ValueError(
- "Could not connect to NebulaGraph database. "
- "Please ensure that the username and password are correct"
- )
- except RuntimeError as e:
- raise ValueError(f"Error initializing session pool. Error: {e}")
-
- return session_pool
-
- def __del__(self) -> None:
- try:
- self.session_pool.close()
- except Exception as e:
- logger.warning(f"Could not close session pool. Error: {e}")
-
- @property
- def get_schema(self) -> str:
- """Returns the schema of the NebulaGraph database"""
- return self.schema
-
- def execute(self, query: str, params: Optional[dict] = None, retry: int = 0) -> Any:
- """Query NebulaGraph database."""
- from nebula3.Exception import IOErrorException, NoValidSessionException
- from nebula3.fbthrift.transport.TTransport import TTransportException
-
- params = params or {}
- try:
- result = self.session_pool.execute_parameter(query, params)
- if not result.is_succeeded():
- logger.warning(
- f"Error executing query to NebulaGraph. "
- f"Error: {result.error_msg()}\n"
- f"Query: {query} \n"
- )
- return result
-
- except NoValidSessionException:
- logger.warning(
- f"No valid session found in session pool. "
- f"Please consider increasing the session pool size. "
- f"Current size: {self.session_pool_size}"
- )
- raise ValueError(
- f"No valid session found in session pool. "
- f"Please consider increasing the session pool size. "
- f"Current size: {self.session_pool_size}"
- )
-
- except RuntimeError as e:
- if retry < RETRY_TIMES:
- retry += 1
- logger.warning(
- f"Error executing query to NebulaGraph. "
- f"Retrying ({retry}/{RETRY_TIMES})...\n"
- f"query: {query} \n"
- f"Error: {e}"
- )
- return self.execute(query, params, retry)
- else:
- raise ValueError(f"Error executing query to NebulaGraph. Error: {e}")
-
- except (TTransportException, IOErrorException):
- # connection issue, try to recreate session pool
- if retry < RETRY_TIMES:
- retry += 1
- logger.warning(
- f"Connection issue with NebulaGraph. "
- f"Retrying ({retry}/{RETRY_TIMES})...\n to recreate session pool"
- )
- self.session_pool = self._get_session_pool()
- return self.execute(query, params, retry)
-
- def refresh_schema(self) -> None:
- """
- Refreshes the NebulaGraph schema information.
- """
- tags_schema, edge_types_schema, relationships = [], [], []
- for tag in self.execute("SHOW TAGS").column_values("Name"):
- tag_name = tag.cast()
- tag_schema = {"tag": tag_name, "properties": []}
- r = self.execute(f"DESCRIBE TAG `{tag_name}`")
- props, types = r.column_values("Field"), r.column_values("Type")
- for i in range(r.row_size()):
- tag_schema["properties"].append((props[i].cast(), types[i].cast()))
- tags_schema.append(tag_schema)
- for edge_type in self.execute("SHOW EDGES").column_values("Name"):
- edge_type_name = edge_type.cast()
- edge_schema = {"edge": edge_type_name, "properties": []}
- r = self.execute(f"DESCRIBE EDGE `{edge_type_name}`")
- props, types = r.column_values("Field"), r.column_values("Type")
- for i in range(r.row_size()):
- edge_schema["properties"].append((props[i].cast(), types[i].cast()))
- edge_types_schema.append(edge_schema)
-
- # build relationships types
- r = self.execute(
- rel_query.substitute(edge_type=edge_type_name)
- ).column_values("rels")
- if len(r) > 0:
- relationships.append(r[0].cast())
-
- self.schema = (
- f"Node properties: {tags_schema}\n"
- f"Edge properties: {edge_types_schema}\n"
- f"Relationships: {relationships}\n"
- )
-
- def query(self, query: str, retry: int = 0) -> Dict[str, Any]:
- result = self.execute(query, retry=retry)
- columns = result.keys()
- d: Dict[str, list] = {}
- for col_num in range(result.col_size()):
- col_name = columns[col_num]
- col_list = result.column_values(col_name)
- d[col_name] = [x.cast() for x in col_list]
- return d
diff --git a/libs/community/langchain_community/graphs/neo4j_graph.py b/libs/community/langchain_community/graphs/neo4j_graph.py
deleted file mode 100644
index 7ce5f2d7e2..0000000000
--- a/libs/community/langchain_community/graphs/neo4j_graph.py
+++ /dev/null
@@ -1,848 +0,0 @@
-from hashlib import md5
-from typing import Any, Dict, List, Optional
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.utils import get_from_dict_or_env
-
-from langchain_community.graphs.graph_document import GraphDocument
-from langchain_community.graphs.graph_store import GraphStore
-
-BASE_ENTITY_LABEL = "__Entity__"
-EXCLUDED_LABELS = ["_Bloom_Perspective_", "_Bloom_Scene_"]
-EXCLUDED_RELS = ["_Bloom_HAS_SCENE_"]
-EXHAUSTIVE_SEARCH_LIMIT = 10000
-LIST_LIMIT = 128
-# Threshold for returning all available prop values in graph schema
-DISTINCT_VALUE_LIMIT = 10
-
-node_properties_query = """
-CALL apoc.meta.data()
-YIELD label, other, elementType, type, property
-WHERE NOT type = "RELATIONSHIP" AND elementType = "node"
- AND NOT label IN $EXCLUDED_LABELS
-WITH label AS nodeLabels, collect({property:property, type:type}) AS properties
-RETURN {labels: nodeLabels, properties: properties} AS output
-
-"""
-
-rel_properties_query = """
-CALL apoc.meta.data()
-YIELD label, other, elementType, type, property
-WHERE NOT type = "RELATIONSHIP" AND elementType = "relationship"
- AND NOT label in $EXCLUDED_LABELS
-WITH label AS nodeLabels, collect({property:property, type:type}) AS properties
-RETURN {type: nodeLabels, properties: properties} AS output
-"""
-
-rel_query = """
-CALL apoc.meta.data()
-YIELD label, other, elementType, type, property
-WHERE type = "RELATIONSHIP" AND elementType = "node"
-UNWIND other AS other_node
-WITH * WHERE NOT label IN $EXCLUDED_LABELS
- AND NOT other_node IN $EXCLUDED_LABELS
-RETURN {start: label, type: property, end: toString(other_node)} AS output
-"""
-
-include_docs_query = (
- "MERGE (d:Document {id:$document.metadata.id}) "
- "SET d.text = $document.page_content "
- "SET d += $document.metadata "
- "WITH d "
-)
-
-
-@deprecated(
- since="0.3.8",
- removal="1.0",
- alternative_import="langchain_neo4j.graphs.neo4j_graph.clean_string_values",
-)
-def clean_string_values(text: str) -> str:
- """Clean string values for schema.
-
- Cleans the input text by replacing newline and carriage return characters.
-
- Args:
- text (str): The input text to clean.
-
- Returns:
- str: The cleaned text.
- """
- return text.replace("\n", " ").replace("\r", " ")
-
-
-@deprecated(
- since="0.3.8",
- removal="1.0",
- alternative_import="langchain_neo4j.graphs.neo4j_graph.value_sanitize",
-)
-def value_sanitize(d: Any) -> Any:
- """Sanitize the input dictionary or list.
-
- Sanitizes the input by removing embedding-like values,
- lists with more than 128 elements, that are mostly irrelevant for
- generating answers in a LLM context. These properties, if left in
- results, can occupy significant context space and detract from
- the LLM's performance by introducing unnecessary noise and cost.
-
- Args:
- d (Any): The input dictionary or list to sanitize.
-
- Returns:
- Any: The sanitized dictionary or list.
- """
- if isinstance(d, dict):
- new_dict = {}
- for key, value in d.items():
- if isinstance(value, dict):
- sanitized_value = value_sanitize(value)
- if (
- sanitized_value is not None
- ): # Check if the sanitized value is not None
- new_dict[key] = sanitized_value
- elif isinstance(value, list):
- if len(value) < LIST_LIMIT:
- sanitized_value = value_sanitize(value)
- if (
- sanitized_value is not None
- ): # Check if the sanitized value is not None
- new_dict[key] = sanitized_value
- # Do not include the key if the list is oversized
- else:
- new_dict[key] = value
- return new_dict
- elif isinstance(d, list):
- if len(d) < LIST_LIMIT:
- return [
- value_sanitize(item) for item in d if value_sanitize(item) is not None
- ]
- else:
- return None
- else:
- return d
-
-
-@deprecated(
- since="0.3.8",
- removal="1.0",
- alternative_import="langchain_neo4j.graphs.neo4j_graph._get_node_import_query",
-)
-def _get_node_import_query(baseEntityLabel: bool, include_source: bool) -> str:
- if baseEntityLabel:
- return (
- f"{include_docs_query if include_source else ''}"
- "UNWIND $data AS row "
- f"MERGE (source:`{BASE_ENTITY_LABEL}` {{id: row.id}}) "
- "SET source += row.properties "
- f"{'MERGE (d)-[:MENTIONS]->(source) ' if include_source else ''}"
- "WITH source, row "
- "CALL apoc.create.addLabels( source, [row.type] ) YIELD node "
- "RETURN distinct 'done' AS result"
- )
- else:
- return (
- f"{include_docs_query if include_source else ''}"
- "UNWIND $data AS row "
- "CALL apoc.merge.node([row.type], {id: row.id}, "
- "row.properties, {}) YIELD node "
- f"{'MERGE (d)-[:MENTIONS]->(node) ' if include_source else ''}"
- "RETURN distinct 'done' AS result"
- )
-
-
-@deprecated(
- since="0.3.8",
- removal="1.0",
- alternative_import="langchain_neo4j.graphs.neo4j_graph._get_rel_import_query",
-)
-def _get_rel_import_query(baseEntityLabel: bool) -> str:
- if baseEntityLabel:
- return (
- "UNWIND $data AS row "
- f"MERGE (source:`{BASE_ENTITY_LABEL}` {{id: row.source}}) "
- f"MERGE (target:`{BASE_ENTITY_LABEL}` {{id: row.target}}) "
- "WITH source, target, row "
- "CALL apoc.merge.relationship(source, row.type, "
- "{}, row.properties, target) YIELD rel "
- "RETURN distinct 'done'"
- )
- else:
- return (
- "UNWIND $data AS row "
- "CALL apoc.merge.node([row.source_label], {id: row.source},"
- "{}, {}) YIELD node as source "
- "CALL apoc.merge.node([row.target_label], {id: row.target},"
- "{}, {}) YIELD node as target "
- "CALL apoc.merge.relationship(source, row.type, "
- "{}, row.properties, target) YIELD rel "
- "RETURN distinct 'done'"
- )
-
-
-@deprecated(
- since="0.3.8",
- removal="1.0",
- alternative_import="langchain_neo4j.graphs.neo4j_graph._format_schema",
-)
-def _format_schema(schema: Dict, is_enhanced: bool) -> str:
- formatted_node_props = []
- formatted_rel_props = []
- if is_enhanced:
- # Enhanced formatting for nodes
- for node_type, properties in schema["node_props"].items():
- formatted_node_props.append(f"- **{node_type}**")
- for prop in properties:
- example = ""
- if prop["type"] == "STRING" and prop.get("values"):
- if prop.get("distinct_count", 11) > DISTINCT_VALUE_LIMIT:
- example = (
- f'Example: "{clean_string_values(prop["values"][0])}"'
- if prop["values"]
- else ""
- )
- else: # If less than 10 possible values return all
- example = (
- (
- "Available options: "
- f"{[clean_string_values(el) for el in prop['values']]}"
- )
- if prop["values"]
- else ""
- )
-
- elif prop["type"] in [
- "INTEGER",
- "FLOAT",
- "DATE",
- "DATE_TIME",
- "LOCAL_DATE_TIME",
- ]:
- if prop.get("min") is not None:
- example = f"Min: {prop['min']}, Max: {prop['max']}"
- else:
- example = (
- f'Example: "{prop["values"][0]}"'
- if prop.get("values")
- else ""
- )
- elif prop["type"] == "LIST":
- # Skip embeddings
- if not prop.get("min_size") or prop["min_size"] > LIST_LIMIT:
- continue
- example = (
- f"Min Size: {prop['min_size']}, Max Size: {prop['max_size']}"
- )
- formatted_node_props.append(
- f" - `{prop['property']}`: {prop['type']} {example}"
- )
-
- # Enhanced formatting for relationships
- for rel_type, properties in schema["rel_props"].items():
- formatted_rel_props.append(f"- **{rel_type}**")
- for prop in properties:
- example = ""
- if prop["type"] == "STRING":
- if prop.get("distinct_count", 11) > DISTINCT_VALUE_LIMIT:
- example = (
- f'Example: "{clean_string_values(prop["values"][0])}"'
- if prop["values"]
- else ""
- )
- else: # If less than 10 possible values return all
- example = (
- (
- "Available options: "
- f"{[clean_string_values(el) for el in prop['values']]}"
- )
- if prop["values"]
- else ""
- )
- elif prop["type"] in [
- "INTEGER",
- "FLOAT",
- "DATE",
- "DATE_TIME",
- "LOCAL_DATE_TIME",
- ]:
- if prop.get("min"): # If we have min/max
- example = f"Min: {prop['min']}, Max: {prop['max']}"
- else: # return a single value
- example = (
- f'Example: "{prop["values"][0]}"' if prop["values"] else ""
- )
- elif prop["type"] == "LIST":
- # Skip embeddings
- if not prop.get("min_size") or prop["min_size"] > LIST_LIMIT:
- continue
- example = (
- f"Min Size: {prop['min_size']}, Max Size: {prop['max_size']}"
- )
- formatted_rel_props.append(
- f" - `{prop['property']}: {prop['type']}` {example}"
- )
- else:
- # Format node properties
- for label, props in schema["node_props"].items():
- props_str = ", ".join(
- [f"{prop['property']}: {prop['type']}" for prop in props]
- )
- formatted_node_props.append(f"{label} {{{props_str}}}")
-
- # Format relationship properties using structured_schema
- for type, props in schema["rel_props"].items():
- props_str = ", ".join(
- [f"{prop['property']}: {prop['type']}" for prop in props]
- )
- formatted_rel_props.append(f"{type} {{{props_str}}}")
-
- # Format relationships
- formatted_rels = [
- f"(:{el['start']})-[:{el['type']}]->(:{el['end']})"
- for el in schema["relationships"]
- ]
-
- return "\n".join(
- [
- "Node properties:",
- "\n".join(formatted_node_props),
- "Relationship properties:",
- "\n".join(formatted_rel_props),
- "The relationships:",
- "\n".join(formatted_rels),
- ]
- )
-
-
-@deprecated(
- since="0.3.8",
- removal="1.0",
- alternative_import="langchain_neo4j.graphs.neo4j_graph._remove_backticks",
-)
-def _remove_backticks(text: str) -> str:
- return text.replace("`", "")
-
-
-@deprecated(
- since="0.3.8",
- removal="1.0",
- alternative_import="langchain_neo4j.Neo4jGraph",
-)
-class Neo4jGraph(GraphStore):
- """Neo4j database wrapper for various graph operations.
-
- Parameters:
- url (Optional[str]): The URL of the Neo4j database server.
- username (Optional[str]): The username for database authentication.
- password (Optional[str]): The password for database authentication.
- database (str): The name of the database to connect to. Default is 'neo4j'.
- timeout (Optional[float]): The timeout for transactions in seconds.
- Useful for terminating long-running queries.
- By default, there is no timeout set.
- sanitize (bool): A flag to indicate whether to remove lists with
- more than 128 elements from results. Useful for removing
- embedding-like properties from database responses. Default is False.
- refresh_schema (bool): A flag whether to refresh schema information
- at initialization. Default is True.
- enhanced_schema (bool): A flag whether to scan the database for
- example values and use them in the graph schema. Default is False.
- driver_config (Dict): Configuration passed to Neo4j Driver.
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- def __init__(
- self,
- url: Optional[str] = None,
- username: Optional[str] = None,
- password: Optional[str] = None,
- database: Optional[str] = None,
- timeout: Optional[float] = None,
- sanitize: bool = False,
- refresh_schema: bool = True,
- *,
- driver_config: Optional[Dict] = None,
- enhanced_schema: bool = False,
- ) -> None:
- """Create a new Neo4j graph wrapper instance."""
- try:
- import neo4j
- except ImportError:
- raise ImportError(
- "Could not import neo4j python package. "
- "Please install it with `pip install neo4j`."
- )
-
- url = get_from_dict_or_env({"url": url}, "url", "NEO4J_URI")
- # if username and password are "", assume Neo4j auth is disabled
- if username == "" and password == "":
- auth = None
- else:
- username = get_from_dict_or_env(
- {"username": username},
- "username",
- "NEO4J_USERNAME",
- )
- password = get_from_dict_or_env(
- {"password": password},
- "password",
- "NEO4J_PASSWORD",
- )
- auth = (username, password)
- database = get_from_dict_or_env(
- {"database": database}, "database", "NEO4J_DATABASE", "neo4j"
- )
-
- self._driver = neo4j.GraphDatabase.driver(
- url, auth=auth, **(driver_config or {})
- )
- self._database = database
- self.timeout = timeout
- self.sanitize = sanitize
- self._enhanced_schema = enhanced_schema
- self.schema: str = ""
- self.structured_schema: Dict[str, Any] = {}
- # Verify connection
- try:
- self._driver.verify_connectivity()
- except neo4j.exceptions.ServiceUnavailable:
- raise ValueError(
- "Could not connect to Neo4j database. "
- "Please ensure that the url is correct"
- )
- except neo4j.exceptions.AuthError:
- raise ValueError(
- "Could not connect to Neo4j database. "
- "Please ensure that the username and password are correct"
- )
- # Set schema
- if refresh_schema:
- try:
- self.refresh_schema()
- except neo4j.exceptions.ClientError as e:
- if e.code == "Neo.ClientError.Procedure.ProcedureNotFound":
- raise ValueError(
- "Could not use APOC procedures. "
- "Please ensure the APOC plugin is installed in Neo4j and that "
- "'apoc.meta.data()' is allowed in Neo4j configuration "
- )
- raise e
-
- @property
- def get_schema(self) -> str:
- """Returns the schema of the Graph"""
- return self.schema
-
- @property
- def get_structured_schema(self) -> Dict[str, Any]:
- """Returns the structured schema of the Graph"""
- return self.structured_schema
-
- def query(
- self,
- query: str,
- params: dict = {},
- ) -> List[Dict[str, Any]]:
- """Query Neo4j database.
-
- Args:
- query (str): The Cypher query to execute.
- params (dict): The parameters to pass to the query.
-
- Returns:
- List[Dict[str, Any]]: The list of dictionaries containing the query results.
- """
- from neo4j import Query
- from neo4j.exceptions import Neo4jError
-
- try:
- data, _, _ = self._driver.execute_query(
- Query(text=query, timeout=self.timeout),
- database_=self._database,
- parameters_=params,
- )
- json_data = [r.data() for r in data]
- if self.sanitize:
- json_data = [value_sanitize(el) for el in json_data]
- return json_data
- except Neo4jError as e:
- if not (
- (
- ( # isCallInTransactionError
- e.code == "Neo.DatabaseError.Statement.ExecutionFailed"
- or e.code
- == "Neo.DatabaseError.Transaction.TransactionStartFailed"
- )
- and "in an implicit transaction" in e.message
- )
- or ( # isPeriodicCommitError
- e.code == "Neo.ClientError.Statement.SemanticError"
- and (
- "in an open transaction is not possible" in e.message
- or "tried to execute in an explicit transaction" in e.message
- )
- )
- ):
- raise
- # fallback to allow implicit transactions
- with self._driver.session(database=self._database) as session:
- data = session.run(Query(text=query, timeout=self.timeout), params)
- json_data = [r.data() for r in data]
- if self.sanitize:
- json_data = [value_sanitize(el) for el in json_data]
- return json_data
-
- def refresh_schema(self) -> None:
- """
- Refreshes the Neo4j graph schema information.
- """
- from neo4j.exceptions import ClientError, CypherTypeError
-
- node_properties = [
- el["output"]
- for el in self.query(
- node_properties_query,
- params={"EXCLUDED_LABELS": EXCLUDED_LABELS + [BASE_ENTITY_LABEL]},
- )
- ]
- rel_properties = [
- el["output"]
- for el in self.query(
- rel_properties_query, params={"EXCLUDED_LABELS": EXCLUDED_RELS}
- )
- ]
- relationships = [
- el["output"]
- for el in self.query(
- rel_query,
- params={"EXCLUDED_LABELS": EXCLUDED_LABELS + [BASE_ENTITY_LABEL]},
- )
- ]
-
- # Get constraints & indexes
- try:
- constraint = self.query("SHOW CONSTRAINTS")
- index = self.query(
- "CALL apoc.schema.nodes() YIELD label, properties, type, size, "
- "valuesSelectivity WHERE type = 'RANGE' RETURN *, "
- "size * valuesSelectivity as distinctValues"
- )
- except (
- ClientError
- ): # Read-only user might not have access to schema information
- constraint = []
- index = []
-
- self.structured_schema = {
- "node_props": {el["labels"]: el["properties"] for el in node_properties},
- "rel_props": {el["type"]: el["properties"] for el in rel_properties},
- "relationships": relationships,
- "metadata": {"constraint": constraint, "index": index},
- }
- if self._enhanced_schema:
- schema_counts = self.query(
- "CALL apoc.meta.graphSample() YIELD nodes, relationships "
- "RETURN nodes, [rel in relationships | {name:apoc.any.property"
- "(rel, 'type'), count: apoc.any.property(rel, 'count')}]"
- " AS relationships"
- )
- # Update node info
- for node in schema_counts[0]["nodes"]:
- # Skip bloom labels
- if node["name"] in EXCLUDED_LABELS:
- continue
- node_props = self.structured_schema["node_props"].get(node["name"])
- if not node_props: # The node has no properties
- continue
- enhanced_cypher = self._enhanced_schema_cypher(
- node["name"], node_props, node["count"] < EXHAUSTIVE_SEARCH_LIMIT
- )
- # Due to schema-flexible nature of neo4j errors can happen
- try:
- enhanced_info = self.query(enhanced_cypher)[0]["output"]
- for prop in node_props:
- if prop["property"] in enhanced_info:
- prop.update(enhanced_info[prop["property"]])
- except CypherTypeError:
- continue
- # Update rel info
- for rel in schema_counts[0]["relationships"]:
- # Skip bloom labels
- if rel["name"] in EXCLUDED_RELS:
- continue
- rel_props = self.structured_schema["rel_props"].get(rel["name"])
- if not rel_props: # The rel has no properties
- continue
- enhanced_cypher = self._enhanced_schema_cypher(
- rel["name"],
- rel_props,
- rel["count"] < EXHAUSTIVE_SEARCH_LIMIT,
- is_relationship=True,
- )
- try:
- enhanced_info = self.query(enhanced_cypher)[0]["output"]
- for prop in rel_props:
- if prop["property"] in enhanced_info:
- prop.update(enhanced_info[prop["property"]])
- # Due to schema-flexible nature of neo4j errors can happen
- except CypherTypeError:
- continue
-
- schema = _format_schema(self.structured_schema, self._enhanced_schema)
-
- self.schema = schema
-
- def add_graph_documents(
- self,
- graph_documents: List[GraphDocument],
- include_source: bool = False,
- baseEntityLabel: bool = False,
- ) -> None:
- """
- This method constructs nodes and relationships in the graph based on the
- provided GraphDocument objects.
-
- Parameters:
- - graph_documents (List[GraphDocument]): A list of GraphDocument objects
- that contain the nodes and relationships to be added to the graph. Each
- GraphDocument should encapsulate the structure of part of the graph,
- including nodes, relationships, and the source document information.
- - include_source (bool, optional): If True, stores the source document
- and links it to nodes in the graph using the MENTIONS relationship.
- This is useful for tracing back the origin of data. Merges source
- documents based on the `id` property from the source document metadata
- if available; otherwise it calculates the MD5 hash of `page_content`
- for merging process. Defaults to False.
- - baseEntityLabel (bool, optional): If True, each newly created node
- gets a secondary __Entity__ label, which is indexed and improves import
- speed and performance. Defaults to False.
- """
- if baseEntityLabel: # Check if constraint already exists
- constraint_exists = any(
- [
- el["labelsOrTypes"] == [BASE_ENTITY_LABEL]
- and el["properties"] == ["id"]
- for el in self.structured_schema.get("metadata", {}).get(
- "constraint", []
- )
- ]
- )
-
- if not constraint_exists:
- # Create constraint
- self.query(
- f"CREATE CONSTRAINT IF NOT EXISTS FOR (b:{BASE_ENTITY_LABEL}) "
- "REQUIRE b.id IS UNIQUE;"
- )
- self.refresh_schema() # Refresh constraint information
-
- node_import_query = _get_node_import_query(baseEntityLabel, include_source)
- rel_import_query = _get_rel_import_query(baseEntityLabel)
- for document in graph_documents:
- if not document.source.metadata.get("id"):
- document.source.metadata["id"] = md5(
- document.source.page_content.encode("utf-8")
- ).hexdigest()
-
- # Remove backticks from node types
- for node in document.nodes:
- node.type = _remove_backticks(node.type)
- # Import nodes
- self.query(
- node_import_query,
- {
- "data": [el.__dict__ for el in document.nodes],
- "document": document.source.__dict__,
- },
- )
- # Import relationships
- self.query(
- rel_import_query,
- {
- "data": [
- {
- "source": el.source.id,
- "source_label": _remove_backticks(el.source.type),
- "target": el.target.id,
- "target_label": _remove_backticks(el.target.type),
- "type": _remove_backticks(
- el.type.replace(" ", "_").upper()
- ),
- "properties": el.properties,
- }
- for el in document.relationships
- ]
- },
- )
-
- def _enhanced_schema_cypher(
- self,
- label_or_type: str,
- properties: List[Dict[str, Any]],
- exhaustive: bool,
- is_relationship: bool = False,
- ) -> str:
- if is_relationship:
- match_clause = f"MATCH ()-[n:`{label_or_type}`]->()"
- else:
- match_clause = f"MATCH (n:`{label_or_type}`)"
-
- with_clauses = []
- return_clauses = []
- output_dict = {}
- if exhaustive:
- for prop in properties:
- prop_name = prop["property"]
- prop_type = prop["type"]
- if prop_type == "STRING":
- with_clauses.append(
- (
- f"collect(distinct substring(toString(n.`{prop_name}`)"
- f", 0, 50)) AS `{prop_name}_values`"
- )
- )
- return_clauses.append(
- (
- f"values:`{prop_name}_values`[..{DISTINCT_VALUE_LIMIT}],"
- f" distinct_count: size(`{prop_name}_values`)"
- )
- )
- elif prop_type in [
- "INTEGER",
- "FLOAT",
- "DATE",
- "DATE_TIME",
- "LOCAL_DATE_TIME",
- ]:
- with_clauses.append(f"min(n.`{prop_name}`) AS `{prop_name}_min`")
- with_clauses.append(f"max(n.`{prop_name}`) AS `{prop_name}_max`")
- with_clauses.append(
- f"count(distinct n.`{prop_name}`) AS `{prop_name}_distinct`"
- )
- return_clauses.append(
- (
- f"min: toString(`{prop_name}_min`), "
- f"max: toString(`{prop_name}_max`), "
- f"distinct_count: `{prop_name}_distinct`"
- )
- )
- elif prop_type == "LIST":
- with_clauses.append(
- (
- f"min(size(n.`{prop_name}`)) AS `{prop_name}_size_min`, "
- f"max(size(n.`{prop_name}`)) AS `{prop_name}_size_max`"
- )
- )
- return_clauses.append(
- f"min_size: `{prop_name}_size_min`, "
- f"max_size: `{prop_name}_size_max`"
- )
- elif prop_type in ["BOOLEAN", "POINT", "DURATION"]:
- continue
- output_dict[prop_name] = "{" + return_clauses.pop() + "}"
- else:
- # Just sample 5 random nodes
- match_clause += " WITH n LIMIT 5"
- for prop in properties:
- prop_name = prop["property"]
- prop_type = prop["type"]
-
- # Check if indexed property, we can still do exhaustive
- prop_index = [
- el
- for el in self.structured_schema["metadata"]["index"]
- if el["label"] == label_or_type
- and el["properties"] == [prop_name]
- and el["type"] == "RANGE"
- ]
- if prop_type == "STRING":
- if (
- prop_index
- and prop_index[0].get("size") > 0
- and prop_index[0].get("distinctValues") <= DISTINCT_VALUE_LIMIT
- ):
- distinct_values = self.query(
- f"CALL apoc.schema.properties.distinct("
- f"'{label_or_type}', '{prop_name}') YIELD value"
- )[0]["value"]
- return_clauses.append(
- (
- f"values: {distinct_values},"
- f" distinct_count: {len(distinct_values)}"
- )
- )
- else:
- with_clauses.append(
- (
- f"collect(distinct substring(toString(n.`{prop_name}`)"
- f", 0, 50)) AS `{prop_name}_values`"
- )
- )
- return_clauses.append(f"values: `{prop_name}_values`")
- elif prop_type in [
- "INTEGER",
- "FLOAT",
- "DATE",
- "DATE_TIME",
- "LOCAL_DATE_TIME",
- ]:
- if not prop_index:
- with_clauses.append(
- f"collect(distinct toString(n.`{prop_name}`)) "
- f"AS `{prop_name}_values`"
- )
- return_clauses.append(f"values: `{prop_name}_values`")
- else:
- with_clauses.append(
- f"min(n.`{prop_name}`) AS `{prop_name}_min`"
- )
- with_clauses.append(
- f"max(n.`{prop_name}`) AS `{prop_name}_max`"
- )
- with_clauses.append(
- f"count(distinct n.`{prop_name}`) AS `{prop_name}_distinct`"
- )
- return_clauses.append(
- (
- f"min: toString(`{prop_name}_min`), "
- f"max: toString(`{prop_name}_max`), "
- f"distinct_count: `{prop_name}_distinct`"
- )
- )
-
- elif prop_type == "LIST":
- with_clauses.append(
- (
- f"min(size(n.`{prop_name}`)) AS `{prop_name}_size_min`, "
- f"max(size(n.`{prop_name}`)) AS `{prop_name}_size_max`"
- )
- )
- return_clauses.append(
- (
- f"min_size: `{prop_name}_size_min`, "
- f"max_size: `{prop_name}_size_max`"
- )
- )
- elif prop_type in ["BOOLEAN", "POINT", "DURATION"]:
- continue
-
- output_dict[prop_name] = "{" + return_clauses.pop() + "}"
-
- with_clause = "WITH " + ",\n ".join(with_clauses)
- return_clause = (
- "RETURN {"
- + ", ".join(f"`{k}`: {v}" for k, v in output_dict.items())
- + "} AS output"
- )
-
- # Combine all parts of the Cypher query
- cypher_query = "\n".join([match_clause, with_clause, return_clause])
- return cypher_query
diff --git a/libs/community/langchain_community/graphs/neptune_graph.py b/libs/community/langchain_community/graphs/neptune_graph.py
deleted file mode 100644
index d71fb07e38..0000000000
--- a/libs/community/langchain_community/graphs/neptune_graph.py
+++ /dev/null
@@ -1,426 +0,0 @@
-import json
-from abc import ABC, abstractmethod
-from typing import Any, Dict, List, Optional, Tuple, Union
-
-from langchain_core._api.deprecation import deprecated
-
-
-class NeptuneQueryException(Exception):
- """Exception for the Neptune queries."""
-
- def __init__(self, exception: Union[str, Dict]):
- if isinstance(exception, dict):
- self.message = exception["message"] if "message" in exception else "unknown"
- self.details = exception["details"] if "details" in exception else "unknown"
- else:
- self.message = exception
- self.details = "unknown"
-
- def get_message(self) -> str:
- return self.message
-
- def get_details(self) -> Any:
- return self.details
-
-
-class BaseNeptuneGraph(ABC):
- """Abstract base class for Neptune."""
-
- @property
- def get_schema(self) -> str:
- """Return the schema of the Neptune database"""
- return self.schema
-
- @abstractmethod
- def query(self, query: str, params: dict = {}) -> dict:
- raise NotImplementedError()
-
- @abstractmethod
- def _get_summary(self) -> Dict:
- raise NotImplementedError()
-
- def _get_labels(self) -> Tuple[List[str], List[str]]:
- """Get node and edge labels from the Neptune statistics summary"""
- summary = self._get_summary()
- n_labels = summary["nodeLabels"]
- e_labels = summary["edgeLabels"]
- return n_labels, e_labels
-
- def _get_triples(self, e_labels: List[str]) -> List[str]:
- triple_query = """
- MATCH (a)-[e:`{e_label}`]->(b)
- WITH a,e,b LIMIT 3000
- RETURN DISTINCT labels(a) AS from, type(e) AS edge, labels(b) AS to
- LIMIT 10
- """
-
- triple_template = "(:`{a}`)-[:`{e}`]->(:`{b}`)"
- triple_schema = []
- for label in e_labels:
- q = triple_query.format(e_label=label)
- data = self.query(q)
- for d in data:
- triple = triple_template.format(
- a=d["from"][0], e=d["edge"], b=d["to"][0]
- )
- triple_schema.append(triple)
-
- return triple_schema
-
- def _get_node_properties(self, n_labels: List[str], types: Dict) -> List:
- node_properties_query = """
- MATCH (a:`{n_label}`)
- RETURN properties(a) AS props
- LIMIT 100
- """
- node_properties = []
- for label in n_labels:
- q = node_properties_query.format(n_label=label)
- data = {"label": label, "properties": self.query(q)}
- s = set({})
- for p in data["properties"]:
- for k, v in p["props"].items():
- s.add((k, types[type(v).__name__]))
-
- np = {
- "properties": [{"property": k, "type": v} for k, v in s],
- "labels": label,
- }
- node_properties.append(np)
-
- return node_properties
-
- def _get_edge_properties(self, e_labels: List[str], types: Dict[str, Any]) -> List:
- edge_properties_query = """
- MATCH ()-[e:`{e_label}`]->()
- RETURN properties(e) AS props
- LIMIT 100
- """
- edge_properties = []
- for label in e_labels:
- q = edge_properties_query.format(e_label=label)
- data = {"label": label, "properties": self.query(q)}
- s = set({})
- for p in data["properties"]:
- for k, v in p["props"].items():
- s.add((k, types[type(v).__name__]))
-
- ep = {
- "type": label,
- "properties": [{"property": k, "type": v} for k, v in s],
- }
- edge_properties.append(ep)
-
- return edge_properties
-
- def _refresh_schema(self) -> None:
- """
- Refreshes the Neptune graph schema information.
- """
-
- types = {
- "str": "STRING",
- "float": "DOUBLE",
- "int": "INTEGER",
- "list": "LIST",
- "dict": "MAP",
- "bool": "BOOLEAN",
- }
- n_labels, e_labels = self._get_labels()
- triple_schema = self._get_triples(e_labels)
- node_properties = self._get_node_properties(n_labels, types)
- edge_properties = self._get_edge_properties(e_labels, types)
-
- self.schema = f"""
- Node properties are the following:
- {node_properties}
- Relationship properties are the following:
- {edge_properties}
- The relationships are the following:
- {triple_schema}
- """
-
-
-@deprecated(
- since="0.3.15",
- removal="1.0",
- alternative_import="langchain_aws.NeptuneAnalyticsGraph",
-)
-class NeptuneAnalyticsGraph(BaseNeptuneGraph):
- """Neptune Analytics wrapper for graph operations.
-
- Parameters:
- client: optional boto3 Neptune client
- credentials_profile_name: optional AWS profile name
- region_name: optional AWS region, e.g., us-west-2
- graph_identifier: the graph identifier for a Neptune Analytics graph
-
- Example:
- .. code-block:: python
-
- graph = NeptuneAnalyticsGraph(
- graph_identifier=''
- )
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- def __init__(
- self,
- graph_identifier: str,
- client: Any = None,
- credentials_profile_name: Optional[str] = None,
- region_name: Optional[str] = None,
- ) -> None:
- """Create a new Neptune Analytics graph wrapper instance."""
-
- try:
- if client is not None:
- self.client = client
- else:
- import boto3
-
- if credentials_profile_name is not None:
- session = boto3.Session(profile_name=credentials_profile_name)
- else:
- # use default credentials
- session = boto3.Session()
-
- self.graph_identifier = graph_identifier
-
- if region_name:
- self.client = session.client(
- "neptune-graph", region_name=region_name
- )
- else:
- self.client = session.client("neptune-graph")
-
- except ImportError:
- raise ImportError(
- "Could not import boto3 python package. "
- "Please install it with `pip install boto3`."
- )
- except Exception as e:
- if type(e).__name__ == "UnknownServiceError":
- raise ImportError(
- "NeptuneGraph requires a boto3 version 1.34.40 or greater."
- "Please install it with `pip install -U boto3`."
- ) from e
- else:
- raise ValueError(
- "Could not load credentials to authenticate with AWS client. "
- "Please check that credentials in the specified "
- "profile name are valid."
- ) from e
-
- try:
- self._refresh_schema()
- except Exception as e:
- raise NeptuneQueryException(
- {
- "message": "Could not get schema for Neptune database",
- "detail": str(e),
- }
- )
-
- def query(self, query: str, params: dict = {}) -> Dict[str, Any]:
- """Query Neptune database."""
- try:
- resp = self.client.execute_query(
- graphIdentifier=self.graph_identifier,
- queryString=query,
- parameters=params,
- language="OPEN_CYPHER",
- )
- return json.loads(resp["payload"].read().decode("UTF-8"))["results"]
- except Exception as e:
- raise NeptuneQueryException(
- {
- "message": "An error occurred while executing the query.",
- "details": str(e),
- }
- )
-
- def _get_summary(self) -> Dict:
- try:
- response = self.client.get_graph_summary(
- graphIdentifier=self.graph_identifier, mode="detailed"
- )
- except Exception as e:
- raise NeptuneQueryException(
- {
- "message": ("Summary API error occurred on Neptune Analytics"),
- "details": str(e),
- }
- )
-
- try:
- summary = response["graphSummary"]
- except Exception:
- raise NeptuneQueryException(
- {
- "message": "Summary API did not return a valid response.",
- "details": response.content.decode(),
- }
- )
- else:
- return summary
-
-
-@deprecated(
- since="0.3.15",
- removal="1.0",
- alternative_import="langchain_aws.NeptuneGraph",
-)
-class NeptuneGraph(BaseNeptuneGraph):
- """Neptune wrapper for graph operations.
-
- Parameters:
- host: endpoint for the database instance
- port: port number for the database instance, default is 8182
- use_https: whether to use secure connection, default is True
- client: optional boto3 Neptune client
- credentials_profile_name: optional AWS profile name
- region_name: optional AWS region, e.g., us-west-2
- sign: optional, whether to sign the request payload, default is True
-
- Example:
- .. code-block:: python
-
- graph = NeptuneGraph(
- host='',
- port=8182
- )
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- def __init__(
- self,
- host: str,
- port: int = 8182,
- use_https: bool = True,
- client: Any = None,
- credentials_profile_name: Optional[str] = None,
- region_name: Optional[str] = None,
- sign: bool = True,
- ) -> None:
- """Create a new Neptune graph wrapper instance."""
-
- try:
- if client is not None:
- self.client = client
- else:
- import boto3
-
- if credentials_profile_name is not None:
- session = boto3.Session(profile_name=credentials_profile_name)
- else:
- # use default credentials
- session = boto3.Session()
-
- client_params = {}
- if region_name:
- client_params["region_name"] = region_name
-
- protocol = "https" if use_https else "http"
-
- client_params["endpoint_url"] = f"{protocol}://{host}:{port}"
-
- if sign:
- self.client = session.client("neptunedata", **client_params)
- else:
- from botocore import UNSIGNED
- from botocore.config import Config
-
- self.client = session.client(
- "neptunedata",
- **client_params,
- config=Config(signature_version=UNSIGNED),
- )
-
- except ImportError:
- raise ImportError(
- "Could not import boto3 python package. "
- "Please install it with `pip install boto3`."
- )
- except Exception as e:
- if type(e).__name__ == "UnknownServiceError":
- raise ImportError(
- "NeptuneGraph requires a boto3 version 1.28.38 or greater."
- "Please install it with `pip install -U boto3`."
- ) from e
- else:
- raise ValueError(
- "Could not load credentials to authenticate with AWS client. "
- "Please check that credentials in the specified "
- "profile name are valid."
- ) from e
-
- try:
- self._refresh_schema()
- except Exception as e:
- raise NeptuneQueryException(
- {
- "message": "Could not get schema for Neptune database",
- "detail": str(e),
- }
- )
-
- def query(self, query: str, params: dict = {}) -> Dict[str, Any]:
- """Query Neptune database."""
- try:
- return self.client.execute_open_cypher_query(openCypherQuery=query)[
- "results"
- ]
- except Exception as e:
- raise NeptuneQueryException(
- {
- "message": "An error occurred while executing the query.",
- "details": str(e),
- }
- )
-
- def _get_summary(self) -> Dict:
- try:
- response = self.client.get_propertygraph_summary()
- except Exception as e:
- raise NeptuneQueryException(
- {
- "message": (
- "Summary API is not available for this instance of Neptune,"
- "ensure the engine version is >=1.2.1.0"
- ),
- "details": str(e),
- }
- )
-
- try:
- summary = response["payload"]["graphSummary"]
- except Exception:
- raise NeptuneQueryException(
- {
- "message": "Summary API did not return a valid response.",
- "details": response.content.decode(),
- }
- )
- else:
- return summary
diff --git a/libs/community/langchain_community/graphs/neptune_rdf_graph.py b/libs/community/langchain_community/graphs/neptune_rdf_graph.py
deleted file mode 100644
index 7f2aefac96..0000000000
--- a/libs/community/langchain_community/graphs/neptune_rdf_graph.py
+++ /dev/null
@@ -1,302 +0,0 @@
-import json
-from types import SimpleNamespace
-from typing import Any, Dict, Optional, Sequence
-
-import requests
-from langchain_core._api.deprecation import deprecated
-
-# Query to find OWL datatype properties
-DTPROP_QUERY = """
-SELECT DISTINCT ?elem
-WHERE {
- ?elem a owl:DatatypeProperty .
-}
-"""
-
-# Query to find OWL object properties
-OPROP_QUERY = """
-SELECT DISTINCT ?elem
-WHERE {
- ?elem a owl:ObjectProperty .
-}
-"""
-
-ELEM_TYPES = {
- "classes": None,
- "rels": None,
- "dtprops": DTPROP_QUERY,
- "oprops": OPROP_QUERY,
-}
-
-
-@deprecated(
- since="0.3.15",
- removal="1.0",
- alternative_import="langchain_aws.NeptuneRdfGraph",
-)
-class NeptuneRdfGraph:
- """Neptune wrapper for RDF graph operations.
-
- Args:
- host: endpoint for the database instance
- port: port number for the database instance, default is 8182
- use_iam_auth: boolean indicating IAM auth is enabled in Neptune cluster
- use_https: whether to use secure connection, default is True
- client: optional boto3 Neptune client
- credentials_profile_name: optional AWS profile name
- region_name: optional AWS region, e.g., us-west-2
- service: optional service name, default is neptunedata
- sign: optional, whether to sign the request payload, default is True
-
- Example:
- .. code-block:: python
-
- graph = NeptuneRdfGraph(
- host=',
- port=
- )
- schema = graph.get_schema()
-
- OR
- graph = NeptuneRdfGraph(
- host=',
- port=
- )
- schema_elem = graph.get_schema_elements()
- #... change schema_elements ...
- graph.load_schema(schema_elem)
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- def __init__(
- self,
- host: str,
- port: int = 8182,
- use_https: bool = True,
- use_iam_auth: bool = False,
- client: Any = None,
- credentials_profile_name: Optional[str] = None,
- region_name: Optional[str] = None,
- service: str = "neptunedata",
- sign: bool = True,
- ) -> None:
- self.use_iam_auth = use_iam_auth
- self.region_name = region_name
- self.query_endpoint = f"https://{host}:{port}/sparql"
-
- try:
- if client is not None:
- self.client = client
- else:
- import boto3
-
- if credentials_profile_name is not None:
- self.session = boto3.Session(profile_name=credentials_profile_name)
- else:
- # use default credentials
- self.session = boto3.Session()
-
- client_params = {}
- if region_name:
- client_params["region_name"] = region_name
-
- protocol = "https" if use_https else "http"
-
- client_params["endpoint_url"] = f"{protocol}://{host}:{port}"
-
- if sign:
- self.client = self.session.client(service, **client_params)
- else:
- from botocore import UNSIGNED
- from botocore.config import Config
-
- self.client = self.session.client(
- service,
- **client_params,
- config=Config(signature_version=UNSIGNED),
- )
-
- except ImportError:
- raise ImportError(
- "Could not import boto3 python package. "
- "Please install it with `pip install boto3`."
- )
- except Exception as e:
- if type(e).__name__ == "UnknownServiceError":
- raise ImportError(
- "NeptuneGraph requires a boto3 version 1.28.38 or greater."
- "Please install it with `pip install -U boto3`."
- ) from e
- else:
- raise ValueError(
- "Could not load credentials to authenticate with AWS client. "
- "Please check that credentials in the specified "
- "profile name are valid."
- ) from e
-
- # Set schema
- self.schema = ""
- self.schema_elements: Dict[str, Any] = {}
- self._refresh_schema()
-
- @property
- def get_schema(self) -> str:
- """
- Returns the schema of the graph database.
- """
- return self.schema
-
- @property
- def get_schema_elements(self) -> Dict[str, Any]:
- return self.schema_elements
-
- def get_summary(self) -> Dict[str, Any]:
- """
- Obtain Neptune statistical summary of classes and predicates in the graph.
- """
- return self.client.get_rdf_graph_summary(mode="detailed")
-
- def query(
- self,
- query: str,
- ) -> Dict[str, Any]:
- """
- Run Neptune query.
- """
- request_data = {"query": query}
- data = request_data
- request_hdr = None
-
- if self.use_iam_auth:
- credentials = self.session.get_credentials()
- credentials = credentials.get_frozen_credentials()
- access_key = credentials.access_key
- secret_key = credentials.secret_key
- service = "neptune-db"
- session_token = credentials.token
- params = None
- creds = SimpleNamespace(
- access_key=access_key,
- secret_key=secret_key,
- token=session_token,
- region=self.region_name,
- )
- from botocore.awsrequest import AWSRequest
-
- request = AWSRequest(
- method="POST", url=self.query_endpoint, data=data, params=params
- )
- from botocore.auth import SigV4Auth
-
- SigV4Auth(creds, service, self.region_name).add_auth(request)
- request.headers["Content-Type"] = "application/x-www-form-urlencoded"
- request_hdr = request.headers
- else:
- request_hdr = {}
- request_hdr["Content-Type"] = "application/x-www-form-urlencoded"
-
- queryres = requests.request(
- method="POST", url=self.query_endpoint, headers=request_hdr, data=data
- )
- json_resp = json.loads(queryres.text)
- return json_resp
-
- def load_schema(self, schema_elements: Dict[str, Any]) -> None:
- """
- Generates and sets schema from schema_elements. Helpful in
- cases where introspected schema needs pruning.
- """
-
- elem_str = {}
- for elem in ELEM_TYPES:
- res_list = []
- for elem_rec in schema_elements[elem]:
- uri = elem_rec["uri"]
- local = elem_rec["local"]
- res_str = f"<{uri}> ({local})"
- res_list.append(res_str)
- elem_str[elem] = ", ".join(res_list)
-
- self.schema = (
- "In the following, each IRI is followed by the local name and "
- "optionally its description in parentheses. \n"
- "The graph supports the following node types:\n"
- f"{elem_str['classes']}\n"
- "The graph supports the following relationships:\n"
- f"{elem_str['rels']}\n"
- "The graph supports the following OWL object properties:\n"
- f"{elem_str['dtprops']}\n"
- "The graph supports the following OWL data properties:\n"
- f"{elem_str['oprops']}"
- )
-
- def _get_local_name(self, iri: str) -> Sequence[str]:
- """
- Split IRI into prefix and local
- """
- if "#" in iri:
- tokens = iri.split("#")
- return [f"{tokens[0]}#", tokens[-1]]
- elif "/" in iri:
- tokens = iri.split("/")
- return [f"{'/'.join(tokens[0 : len(tokens) - 1])}/", tokens[-1]]
- else:
- raise ValueError(f"Unexpected IRI '{iri}', contains neither '#' nor '/'.")
-
- def _refresh_schema(self) -> None:
- """
- Query Neptune to introspect schema.
- """
- self.schema_elements["distinct_prefixes"] = {}
-
- # get summary and build list of classes and rels
- summary = self.get_summary()
- reslist = []
- for c in summary["payload"]["graphSummary"]["classes"]:
- uri = c
- tokens = self._get_local_name(uri)
- elem_record = {"uri": uri, "local": tokens[1]}
- reslist.append(elem_record)
- if tokens[0] not in self.schema_elements["distinct_prefixes"]:
- self.schema_elements["distinct_prefixes"][tokens[0]] = "y"
- self.schema_elements["classes"] = reslist
-
- reslist = []
- for r in summary["payload"]["graphSummary"]["predicates"]:
- for p in r:
- uri = p
- tokens = self._get_local_name(uri)
- elem_record = {"uri": uri, "local": tokens[1]}
- reslist.append(elem_record)
- if tokens[0] not in self.schema_elements["distinct_prefixes"]:
- self.schema_elements["distinct_prefixes"][tokens[0]] = "y"
- self.schema_elements["rels"] = reslist
-
- # get dtprops and oprops too
- for elem in ELEM_TYPES:
- q = ELEM_TYPES.get(elem)
- if not q:
- continue
- items = self.query(q)
- reslist = []
- for r in items["results"]["bindings"]:
- uri = r["elem"]["value"]
- tokens = self._get_local_name(uri)
- elem_record = {"uri": uri, "local": tokens[1]}
- reslist.append(elem_record)
- if tokens[0] not in self.schema_elements["distinct_prefixes"]:
- self.schema_elements["distinct_prefixes"][tokens[0]] = "y"
-
- self.schema_elements[elem] = reslist
-
- self.load_schema(self.schema_elements)
diff --git a/libs/community/langchain_community/graphs/networkx_graph.py b/libs/community/langchain_community/graphs/networkx_graph.py
deleted file mode 100644
index 28e78fcc1a..0000000000
--- a/libs/community/langchain_community/graphs/networkx_graph.py
+++ /dev/null
@@ -1,218 +0,0 @@
-"""Networkx wrapper for graph operations."""
-
-from __future__ import annotations
-
-from typing import Any, List, NamedTuple, Optional, Tuple
-
-KG_TRIPLE_DELIMITER = "<|>"
-
-
-class KnowledgeTriple(NamedTuple):
- """Knowledge triple in the graph."""
-
- subject: str
- predicate: str
- object_: str
-
- @classmethod
- def from_string(cls, triple_string: str) -> "KnowledgeTriple":
- """Create a KnowledgeTriple from a string."""
- subject, predicate, object_ = triple_string.strip().split(", ")
- subject = subject[1:]
- object_ = object_[:-1]
- return cls(subject, predicate, object_)
-
-
-def parse_triples(knowledge_str: str) -> List[KnowledgeTriple]:
- """Parse knowledge triples from the knowledge string."""
- knowledge_str = knowledge_str.strip()
- if not knowledge_str or knowledge_str == "NONE":
- return []
- triple_strs = knowledge_str.split(KG_TRIPLE_DELIMITER)
- results = []
- for triple_str in triple_strs:
- try:
- kg_triple = KnowledgeTriple.from_string(triple_str)
- except ValueError:
- continue
- results.append(kg_triple)
- return results
-
-
-def get_entities(entity_str: str) -> List[str]:
- """Extract entities from entity string."""
- if entity_str.strip() == "NONE":
- return []
- else:
- return [w.strip() for w in entity_str.split(",")]
-
-
-class NetworkxEntityGraph:
- """Networkx wrapper for entity graph operations.
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- def __init__(self, graph: Optional[Any] = None) -> None:
- """Create a new graph."""
- try:
- import networkx as nx
- except ImportError:
- raise ImportError(
- "Could not import networkx python package. "
- "Please install it with `pip install networkx`."
- )
- if graph is not None:
- if not isinstance(graph, nx.DiGraph):
- raise ValueError("Passed in graph is not of correct shape")
- self._graph = graph
- else:
- self._graph = nx.DiGraph()
-
- @classmethod
- def from_gml(cls, gml_path: str) -> NetworkxEntityGraph:
- try:
- import networkx as nx
- except ImportError:
- raise ImportError(
- "Could not import networkx python package. "
- "Please install it with `pip install networkx`."
- )
- graph = nx.read_gml(gml_path)
- return cls(graph)
-
- def add_triple(self, knowledge_triple: KnowledgeTriple) -> None:
- """Add a triple to the graph."""
- # Creates nodes if they don't exist
- # Overwrites existing edges
- if not self._graph.has_node(knowledge_triple.subject):
- self._graph.add_node(knowledge_triple.subject)
- if not self._graph.has_node(knowledge_triple.object_):
- self._graph.add_node(knowledge_triple.object_)
- self._graph.add_edge(
- knowledge_triple.subject,
- knowledge_triple.object_,
- relation=knowledge_triple.predicate,
- )
-
- def delete_triple(self, knowledge_triple: KnowledgeTriple) -> None:
- """Delete a triple from the graph."""
- if self._graph.has_edge(knowledge_triple.subject, knowledge_triple.object_):
- self._graph.remove_edge(knowledge_triple.subject, knowledge_triple.object_)
-
- def get_triples(self) -> List[Tuple[str, str, str]]:
- """Get all triples in the graph."""
- return [(u, v, d["relation"]) for u, v, d in self._graph.edges(data=True)]
-
- def get_entity_knowledge(self, entity: str, depth: int = 1) -> List[str]:
- """Get information about an entity."""
- import networkx as nx
-
- # TODO: Have more information-specific retrieval methods
- if not self._graph.has_node(entity):
- return []
-
- results = []
- for src, sink in nx.dfs_edges(self._graph, entity, depth_limit=depth):
- relation = self._graph[src][sink]["relation"]
- results.append(f"{src} {relation} {sink}")
- return results
-
- def write_to_gml(self, path: str) -> None:
- import networkx as nx
-
- nx.write_gml(self._graph, path)
-
- def clear(self) -> None:
- """Clear the graph."""
- self._graph.clear()
-
- def clear_edges(self) -> None:
- """Clear the graph edges."""
- self._graph.clear_edges()
-
- def add_node(self, node: str) -> None:
- """Add node in the graph."""
- self._graph.add_node(node)
-
- def remove_node(self, node: str) -> None:
- """Remove node from the graph."""
- if self._graph.has_node(node):
- self._graph.remove_node(node)
-
- def has_node(self, node: str) -> bool:
- """Return if graph has the given node."""
- return self._graph.has_node(node)
-
- def remove_edge(self, source_node: str, destination_node: str) -> None:
- """Remove edge from the graph."""
- self._graph.remove_edge(source_node, destination_node)
-
- def has_edge(self, source_node: str, destination_node: str) -> bool:
- """Return if graph has an edge between the given nodes."""
- if self._graph.has_node(source_node) and self._graph.has_node(destination_node):
- return self._graph.has_edge(source_node, destination_node)
- else:
- return False
-
- def get_neighbors(self, node: str) -> List[str]:
- """Return the neighbor nodes of the given node."""
- return self._graph.neighbors(node)
-
- def get_number_of_nodes(self) -> int:
- """Get number of nodes in the graph."""
- return self._graph.number_of_nodes()
-
- def get_topological_sort(self) -> List[str]:
- """Get a list of entity names in the graph sorted by causal dependence."""
- import networkx as nx
-
- return list(nx.topological_sort(self._graph))
-
- def draw_graphviz(self, **kwargs: Any) -> None:
- """
- Provides better drawing
-
- Usage in a jupyter notebook:
-
- >>> from IPython.display import SVG
- >>> self.draw_graphviz_svg(layout="dot", filename="web.svg")
- >>> SVG('web.svg')
- """
- from networkx.drawing.nx_agraph import to_agraph
-
- try:
- import pygraphviz # noqa: F401
-
- except ImportError as e:
- if e.name == "_graphviz":
- """
- >>> e.msg # pygraphviz throws this error
- ImportError: libcgraph.so.6: cannot open shared object file
- """
- raise ImportError(
- "Could not import graphviz debian package. "
- "Please install it with:"
- "`sudo apt-get update`"
- "`sudo apt-get install graphviz graphviz-dev`"
- )
- else:
- raise ImportError(
- "Could not import pygraphviz python package. "
- "Please install it with:"
- "`pip install pygraphviz`."
- )
-
- graph = to_agraph(self._graph) # --> pygraphviz.agraph.AGraph
- # pygraphviz.github.io/documentation/stable/tutorial.html#layout-and-drawing
- graph.layout(prog=kwargs.get("prog", "dot"))
- graph.draw(kwargs.get("path", "graph.svg"))
diff --git a/libs/community/langchain_community/graphs/ontotext_graphdb_graph.py b/libs/community/langchain_community/graphs/ontotext_graphdb_graph.py
deleted file mode 100644
index aa3606cf27..0000000000
--- a/libs/community/langchain_community/graphs/ontotext_graphdb_graph.py
+++ /dev/null
@@ -1,210 +0,0 @@
-from __future__ import annotations
-
-import os
-from typing import (
- TYPE_CHECKING,
- List,
- Optional,
- Union,
-)
-
-if TYPE_CHECKING:
- import rdflib
-
-
-class OntotextGraphDBGraph:
- """Ontotext GraphDB https://graphdb.ontotext.com/ wrapper for graph operations.
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- def __init__(
- self,
- query_endpoint: str,
- query_ontology: Optional[str] = None,
- local_file: Optional[str] = None,
- local_file_format: Optional[str] = None,
- ) -> None:
- """
- Set up the GraphDB wrapper
-
- :param query_endpoint: SPARQL endpoint for queries, read access
-
- If GraphDB is secured,
- set the environment variables 'GRAPHDB_USERNAME' and 'GRAPHDB_PASSWORD'.
-
- :param query_ontology: a `CONSTRUCT` query that is executed
- on the SPARQL endpoint and returns the KG schema statements
- Example:
- 'CONSTRUCT {?s ?p ?o} FROM WHERE {?s ?p ?o}'
- Currently, DESCRIBE queries like
- 'PREFIX onto:
- PREFIX rdfs:
- DESCRIBE ?term WHERE {
- ?term rdfs:isDefinedBy onto:
- }'
- are not supported, because DESCRIBE returns
- the Symmetric Concise Bounded Description (SCBD),
- i.e. also the incoming class links.
- In case of large graphs with a million of instances, this is not efficient.
- Check https://github.com/eclipse-rdf4j/rdf4j/issues/4857
-
- :param local_file: a local RDF ontology file.
- Supported RDF formats:
- Turtle, RDF/XML, JSON-LD, N-Triples, Notation-3, Trig, Trix, N-Quads.
- If the rdf format can't be determined from the file extension,
- pass explicitly the rdf format in `local_file_format` param.
-
- :param local_file_format: Used if the rdf format can't be determined
- from the local file extension.
- One of "json-ld", "xml", "n3", "turtle", "nt", "trig", "nquads", "trix"
-
- Either `query_ontology` or `local_file` should be passed.
- """
-
- if query_ontology and local_file:
- raise ValueError("Both file and query provided. Only one is allowed.")
-
- if not query_ontology and not local_file:
- raise ValueError("Neither file nor query provided. One is required.")
-
- try:
- import rdflib
- from rdflib.plugins.stores import sparqlstore
- except ImportError:
- raise ImportError(
- "Could not import rdflib python package. "
- "Please install it with `pip install rdflib`."
- )
-
- auth = self._get_auth()
- store = sparqlstore.SPARQLStore(auth=auth)
- store.open(query_endpoint)
-
- self.graph = rdflib.Graph(store, identifier=None, bind_namespaces="none")
- self._check_connectivity()
-
- if local_file:
- ontology_schema_graph = self._load_ontology_schema_from_file(
- local_file,
- local_file_format, # type: ignore[arg-type]
- )
- else:
- self._validate_user_query(query_ontology) # type: ignore[arg-type]
- ontology_schema_graph = self._load_ontology_schema_with_query(
- query_ontology # type: ignore[arg-type]
- )
- self.schema = ontology_schema_graph.serialize(format="turtle")
-
- @staticmethod
- def _get_auth() -> Union[tuple, None]:
- """
- Returns the basic authentication configuration
- """
- username = os.environ.get("GRAPHDB_USERNAME", None)
- password = os.environ.get("GRAPHDB_PASSWORD", None)
-
- if username:
- if not password:
- raise ValueError(
- "Environment variable 'GRAPHDB_USERNAME' is set, "
- "but 'GRAPHDB_PASSWORD' is not set."
- )
- else:
- return username, password
- return None
-
- def _check_connectivity(self) -> None:
- """
- Executes a simple `ASK` query to check connectivity
- """
- try:
- self.graph.query("ASK { ?s ?p ?o }")
- except ValueError:
- raise ValueError(
- "Could not query the provided endpoint. "
- "Please, check, if the value of the provided "
- "query_endpoint points to the right repository. "
- "If GraphDB is secured, please, "
- "make sure that the environment variables "
- "'GRAPHDB_USERNAME' and 'GRAPHDB_PASSWORD' are set."
- )
-
- @staticmethod
- def _load_ontology_schema_from_file(local_file: str, local_file_format: str = None): # type: ignore[no-untyped-def, assignment]
- """
- Parse the ontology schema statements from the provided file
- """
- import rdflib
-
- if not os.path.exists(local_file):
- raise FileNotFoundError(f"File {local_file} does not exist.")
- if not os.access(local_file, os.R_OK):
- raise PermissionError(f"Read permission for {local_file} is restricted")
- graph = rdflib.ConjunctiveGraph()
- try:
- graph.parse(local_file, format=local_file_format)
- except Exception as e:
- raise ValueError(f"Invalid file format for {local_file} : ", e)
- return graph
-
- @staticmethod
- def _validate_user_query(query_ontology: str) -> None:
- """
- Validate the query is a valid SPARQL CONSTRUCT query
- """
- from pyparsing import ParseException
- from rdflib.plugins.sparql import prepareQuery
-
- if not isinstance(query_ontology, str):
- raise TypeError("Ontology query must be provided as string.")
- try:
- parsed_query = prepareQuery(query_ontology)
- except ParseException as e:
- raise ValueError("Ontology query is not a valid SPARQL query.", e)
-
- if parsed_query.algebra.name != "ConstructQuery":
- raise ValueError(
- "Invalid query type. Only CONSTRUCT queries are supported."
- )
-
- def _load_ontology_schema_with_query(self, query: str): # type: ignore[no-untyped-def]
- """
- Execute the query for collecting the ontology schema statements
- """
- from rdflib.exceptions import ParserError
-
- try:
- results = self.graph.query(query)
- except ParserError as e:
- raise ValueError(f"Generated SPARQL statement is invalid\n{e}")
-
- return results.graph
-
- @property
- def get_schema(self) -> str:
- """
- Returns the schema of the graph database in turtle format
- """
- return self.schema
-
- def query(
- self,
- query: str,
- ) -> List[rdflib.query.ResultRow]:
- """
- Query the graph.
- """
- from rdflib.query import ResultRow
-
- res = self.graph.query(query)
- return [r for r in res if isinstance(r, ResultRow)]
diff --git a/libs/community/langchain_community/graphs/rdf_graph.py b/libs/community/langchain_community/graphs/rdf_graph.py
deleted file mode 100644
index ca8595c162..0000000000
--- a/libs/community/langchain_community/graphs/rdf_graph.py
+++ /dev/null
@@ -1,307 +0,0 @@
-from __future__ import annotations
-
-from typing import (
- TYPE_CHECKING,
- Dict,
- List,
- Optional,
-)
-
-if TYPE_CHECKING:
- import rdflib
-
-prefixes = {
- "owl": """PREFIX owl: \n""",
- "rdf": """PREFIX rdf: \n""",
- "rdfs": """PREFIX rdfs: \n""",
- "xsd": """PREFIX xsd: \n""",
-}
-
-cls_query_rdf = prefixes["rdfs"] + (
- """SELECT DISTINCT ?cls ?com\n"""
- """WHERE { \n"""
- """ ?instance a ?cls . \n"""
- """ OPTIONAL { ?cls rdfs:comment ?com } \n"""
- """}"""
-)
-
-cls_query_rdfs = prefixes["rdfs"] + (
- """SELECT DISTINCT ?cls ?com\n"""
- """WHERE { \n"""
- """ ?instance a/rdfs:subClassOf* ?cls . \n"""
- """ OPTIONAL { ?cls rdfs:comment ?com } \n"""
- """}"""
-)
-
-cls_query_owl = prefixes["rdfs"] + (
- """SELECT DISTINCT ?cls ?com\n"""
- """WHERE { \n"""
- """ ?instance a/rdfs:subClassOf* ?cls . \n"""
- """ FILTER (isIRI(?cls)) . \n"""
- """ OPTIONAL { ?cls rdfs:comment ?com } \n"""
- """}"""
-)
-
-rel_query_rdf = prefixes["rdfs"] + (
- """SELECT DISTINCT ?rel ?com\n"""
- """WHERE { \n"""
- """ ?subj ?rel ?obj . \n"""
- """ OPTIONAL { ?rel rdfs:comment ?com } \n"""
- """}"""
-)
-
-rel_query_rdfs = (
- prefixes["rdf"]
- + prefixes["rdfs"]
- + (
- """SELECT DISTINCT ?rel ?com\n"""
- """WHERE { \n"""
- """ ?rel a/rdfs:subPropertyOf* rdf:Property . \n"""
- """ OPTIONAL { ?rel rdfs:comment ?com } \n"""
- """}"""
- )
-)
-
-op_query_owl = (
- prefixes["rdfs"]
- + prefixes["owl"]
- + (
- """SELECT DISTINCT ?op ?com\n"""
- """WHERE { \n"""
- """ ?op a/rdfs:subPropertyOf* owl:ObjectProperty . \n"""
- """ OPTIONAL { ?op rdfs:comment ?com } \n"""
- """}"""
- )
-)
-
-dp_query_owl = (
- prefixes["rdfs"]
- + prefixes["owl"]
- + (
- """SELECT DISTINCT ?dp ?com\n"""
- """WHERE { \n"""
- """ ?dp a/rdfs:subPropertyOf* owl:DatatypeProperty . \n"""
- """ OPTIONAL { ?dp rdfs:comment ?com } \n"""
- """}"""
- )
-)
-
-
-class RdfGraph:
- """RDFlib wrapper for graph operations.
-
- Modes:
- * local: Local file - can be queried and changed
- * online: Online file - can only be queried, changes can be stored locally
- * store: Triple store - can be queried and changed if update_endpoint available
- Together with a source file, the serialization should be specified.
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- def __init__(
- self,
- source_file: Optional[str] = None,
- serialization: Optional[str] = "ttl",
- query_endpoint: Optional[str] = None,
- update_endpoint: Optional[str] = None,
- standard: Optional[str] = "rdf",
- local_copy: Optional[str] = None,
- graph_kwargs: Optional[Dict] = None,
- store_kwargs: Optional[Dict] = None,
- ) -> None:
- """
- Set up the RDFlib graph
-
- :param source_file: either a path for a local file or a URL
- :param serialization: serialization of the input
- :param query_endpoint: SPARQL endpoint for queries, read access
- :param update_endpoint: SPARQL endpoint for UPDATE queries, write access
- :param standard: RDF, RDFS, or OWL
- :param local_copy: new local copy for storing changes
- :param graph_kwargs: Additional rdflib.Graph specific kwargs
- that will be used to initialize it,
- if query_endpoint is provided.
- :param store_kwargs: Additional sparqlstore.SPARQLStore specific kwargs
- that will be used to initialize it,
- if query_endpoint is provided.
- """
- self.source_file = source_file
- self.serialization = serialization
- self.query_endpoint = query_endpoint
- self.update_endpoint = update_endpoint
- self.standard = standard
- self.local_copy = local_copy
-
- try:
- import rdflib
- from rdflib.plugins.stores import sparqlstore
- except ImportError:
- raise ImportError(
- "Could not import rdflib python package. "
- "Please install it with `pip install rdflib`."
- )
- if self.standard not in (supported_standards := ("rdf", "rdfs", "owl")):
- raise ValueError(
- f"Invalid standard. Supported standards are: {supported_standards}."
- )
-
- if (
- not source_file
- and not query_endpoint
- or source_file
- and (query_endpoint or update_endpoint)
- ):
- raise ValueError(
- "Could not unambiguously initialize the graph wrapper. "
- "Specify either a file (local or online) via the source_file "
- "or a triple store via the endpoints."
- )
-
- if source_file:
- if source_file.startswith("http"):
- self.mode = "online"
- else:
- self.mode = "local"
- if self.local_copy is None:
- self.local_copy = self.source_file
- self.graph = rdflib.Graph()
- self.graph.parse(source_file, format=self.serialization)
-
- if query_endpoint:
- store_kwargs = store_kwargs or {}
- self.mode = "store"
- if not update_endpoint:
- self._store = sparqlstore.SPARQLStore(**store_kwargs)
- self._store.open(query_endpoint)
- else:
- self._store = sparqlstore.SPARQLUpdateStore(**store_kwargs)
- self._store.open((query_endpoint, update_endpoint))
- graph_kwargs = graph_kwargs or {}
- self.graph = rdflib.Graph(self._store, **graph_kwargs)
-
- # Verify that the graph was loaded
- if not len(self.graph):
- raise AssertionError("The graph is empty.")
-
- # Set schema
- self.schema = ""
- self.load_schema()
-
- @property
- def get_schema(self) -> str:
- """
- Returns the schema of the graph database.
- """
- return self.schema
-
- def query(
- self,
- query: str,
- ) -> List[rdflib.query.ResultRow]:
- """
- Query the graph.
- """
- from rdflib.exceptions import ParserError
- from rdflib.query import ResultRow
-
- try:
- res = self.graph.query(query)
- except ParserError as e:
- raise ValueError(f"Generated SPARQL statement is invalid\n{e}")
- return [r for r in res if isinstance(r, ResultRow)]
-
- def update(
- self,
- query: str,
- ) -> None:
- """
- Update the graph.
- """
- from rdflib.exceptions import ParserError
-
- try:
- self.graph.update(query)
- except ParserError as e:
- raise ValueError(f"Generated SPARQL statement is invalid\n{e}")
- if self.local_copy:
- self.graph.serialize(
- destination=self.local_copy, format=self.local_copy.split(".")[-1]
- )
- else:
- raise ValueError("No target file specified for saving the updated file.")
-
- @staticmethod
- def _get_local_name(iri: str) -> str:
- if "#" in iri:
- local_name = iri.split("#")[-1]
- elif "/" in iri:
- local_name = iri.split("/")[-1]
- else:
- raise ValueError(f"Unexpected IRI '{iri}', contains neither '#' nor '/'.")
- return local_name
-
- def _res_to_str(self, res: rdflib.query.ResultRow, var: str) -> str:
- return (
- "<"
- + str(res[var])
- + "> ("
- + self._get_local_name(res[var])
- + ", "
- + str(res["com"])
- + ")"
- )
-
- def load_schema(self) -> None:
- """
- Load the graph schema information.
- """
-
- def _rdf_s_schema(
- classes: List[rdflib.query.ResultRow],
- relationships: List[rdflib.query.ResultRow],
- ) -> str:
- return (
- f"In the following, each IRI is followed by the local name and "
- f"optionally its description in parentheses. \n"
- f"The RDF graph supports the following node types:\n"
- f"{', '.join([self._res_to_str(r, 'cls') for r in classes])}\n"
- f"The RDF graph supports the following relationships:\n"
- f"{', '.join([self._res_to_str(r, 'rel') for r in relationships])}\n"
- )
-
- if self.standard == "rdf":
- clss = self.query(cls_query_rdf)
- rels = self.query(rel_query_rdf)
- self.schema = _rdf_s_schema(clss, rels)
- elif self.standard == "rdfs":
- clss = self.query(cls_query_rdfs)
- rels = self.query(rel_query_rdfs)
- self.schema = _rdf_s_schema(clss, rels)
- elif self.standard == "owl":
- clss = self.query(cls_query_owl)
- ops = self.query(op_query_owl)
- dps = self.query(dp_query_owl)
- self.schema = (
- f"In the following, each IRI is followed by the local name and "
- f"optionally its description in parentheses. \n"
- f"The OWL graph supports the following node types:\n"
- f"{', '.join([self._res_to_str(r, 'cls') for r in clss])}\n"
- f"The OWL graph supports the following object properties, "
- f"i.e., relationships between objects:\n"
- f"{', '.join([self._res_to_str(r, 'op') for r in ops])}\n"
- f"The OWL graph supports the following data properties, "
- f"i.e., relationships between objects and literals:\n"
- f"{', '.join([self._res_to_str(r, 'dp') for r in dps])}\n"
- )
- else:
- raise ValueError(f"Mode '{self.standard}' is currently not supported.")
diff --git a/libs/community/langchain_community/graphs/tigergraph_graph.py b/libs/community/langchain_community/graphs/tigergraph_graph.py
deleted file mode 100644
index 84b24218ad..0000000000
--- a/libs/community/langchain_community/graphs/tigergraph_graph.py
+++ /dev/null
@@ -1,100 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-from langchain_community.graphs.graph_store import GraphStore
-
-
-class TigerGraph(GraphStore):
- """TigerGraph wrapper for graph operations.
-
- *Security note*: Make sure that the database connection uses credentials
- that are narrowly-scoped to only include necessary permissions.
- Failure to do so may result in data corruption or loss, since the calling
- code may attempt commands that would result in deletion, mutation
- of data if appropriately prompted or reading sensitive data if such
- data is present in the database.
- The best way to guard against such negative outcomes is to (as appropriate)
- limit the permissions granted to the credentials used with this tool.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- def __init__(self, conn: Any) -> None:
- """Create a new TigerGraph graph wrapper instance."""
- self.set_connection(conn)
- self.set_schema()
-
- @property
- def conn(self) -> Any:
- return self._conn
-
- @property
- def schema(self) -> Dict[str, Any]:
- return self._schema
-
- def get_schema(self) -> str: # type: ignore[override]
- if self._schema:
- return str(self._schema)
- else:
- self.set_schema()
- return str(self._schema)
-
- def set_connection(self, conn: Any) -> None:
- try:
- from pyTigerGraph import TigerGraphConnection
- except ImportError:
- raise ImportError(
- "Could not import pyTigerGraph python package. "
- "Please install it with `pip install pyTigerGraph`."
- )
-
- if not isinstance(conn, TigerGraphConnection):
- msg = "**conn** parameter must inherit from TigerGraphConnection"
- raise TypeError(msg)
-
- if conn.ai.nlqs_host is None:
- msg = """**conn** parameter does not have nlqs_host parameter defined.
- Define hostname of NLQS service."""
- raise ConnectionError(msg)
-
- self._conn: TigerGraphConnection = conn
- self.set_schema()
-
- def set_schema(self, schema: Optional[Dict[str, Any]] = None) -> None:
- """
- Set the schema of the TigerGraph Database.
- Auto-generates Schema if **schema** is None.
- """
- self._schema = self.generate_schema() if schema is None else schema
-
- def generate_schema(
- self,
- ) -> Dict[str, List[Dict[str, Any]]]:
- """
- Generates the schema of the TigerGraph Database and returns it
- User can specify a **sample_ratio** (0 to 1) to determine the
- ratio of documents/edges used (in relation to the Collection size)
- to render each Collection Schema.
- """
- return self._conn.getSchema(force=True)
-
- def refresh_schema(self): # type: ignore[no-untyped-def]
- self.generate_schema()
-
- def query(self, query: str) -> Dict[str, Any]: # type: ignore[override]
- """Query the TigerGraph database."""
- answer = self._conn.ai.query(query)
- return answer
-
- def register_query(
- self,
- function_header: str,
- description: str,
- docstring: str,
- param_types: dict = {},
- ) -> List[str]:
- """
- Wrapper function to register a custom GSQL query to the TigerGraph NLQS.
- """
- return self._conn.ai.registerCustomQuery(
- function_header, description, docstring, param_types
- )
diff --git a/libs/community/langchain_community/indexes/__init__.py b/libs/community/langchain_community/indexes/__init__.py
deleted file mode 100644
index 2810a09899..0000000000
--- a/libs/community/langchain_community/indexes/__init__.py
+++ /dev/null
@@ -1,13 +0,0 @@
-"""**Index** is used to avoid writing duplicated content
-into the vectostore and to avoid over-writing content if it's unchanged.
-
-Indexes also :
-
-* Create knowledge graphs from data.
-
-* Support indexing workflows from LangChain data loaders to vectorstores.
-
-Importantly, Index keeps on working even if the content being written is derived
-via a set of transformations from some source content (e.g., indexing children
-documents that were derived from parent documents by chunking.)
-"""
diff --git a/libs/community/langchain_community/indexes/_document_manager.py b/libs/community/langchain_community/indexes/_document_manager.py
deleted file mode 100644
index 45dc2476ff..0000000000
--- a/libs/community/langchain_community/indexes/_document_manager.py
+++ /dev/null
@@ -1,237 +0,0 @@
-from typing import Any, Dict, List, Optional, Sequence
-
-from langchain_community.indexes.base import RecordManager
-
-IMPORT_PYMONGO_ERROR = (
- "Could not import MongoClient. Please install it with `pip install pymongo`."
-)
-IMPORT_MOTOR_ASYNCIO_ERROR = (
- "Could not import AsyncIOMotorClient. Please install it with `pip install motor`."
-)
-
-
-def _import_pymongo() -> Any:
- """Import PyMongo if available, otherwise raise error."""
- try:
- from pymongo import MongoClient
- except ImportError:
- raise ImportError(IMPORT_PYMONGO_ERROR)
- return MongoClient
-
-
-def _get_pymongo_client(mongodb_url: str, **kwargs: Any) -> Any:
- """Get MongoClient for sync operations from the mongodb_url,
- otherwise raise error."""
- try:
- pymongo = _import_pymongo()
- client = pymongo(mongodb_url, **kwargs)
- except ValueError as e:
- raise ImportError(
- f"MongoClient string provided is not in proper format. Got error: {e} "
- )
- return client
-
-
-def _import_motor_asyncio() -> Any:
- """Import Motor if available, otherwise raise error."""
- try:
- from motor.motor_asyncio import AsyncIOMotorClient
- except ImportError:
- raise ImportError(IMPORT_MOTOR_ASYNCIO_ERROR)
- return AsyncIOMotorClient
-
-
-def _get_motor_client(mongodb_url: str, **kwargs: Any) -> Any:
- """Get AsyncIOMotorClient for async operations from the mongodb_url,
- otherwise raise error."""
- try:
- motor = _import_motor_asyncio()
- client = motor(mongodb_url, **kwargs)
- except ValueError as e:
- raise ImportError(
- f"AsyncIOMotorClient string provided is not in proper format. "
- f"Got error: {e} "
- )
- return client
-
-
-class MongoDocumentManager(RecordManager):
- """A MongoDB based implementation of the document manager."""
-
- def __init__(
- self,
- namespace: str,
- *,
- mongodb_url: str,
- db_name: str,
- collection_name: str = "documentMetadata",
- ) -> None:
- """Initialize the MongoDocumentManager.
-
- Args:
- namespace: The namespace associated with this document manager.
- db_name: The name of the database to use.
- collection_name: The name of the collection to use.
- Default is 'documentMetadata'.
- """
- super().__init__(namespace=namespace)
- self.sync_client = _get_pymongo_client(mongodb_url)
- self.sync_db = self.sync_client[db_name]
- self.sync_collection = self.sync_db[collection_name]
- self.async_client = _get_motor_client(mongodb_url)
- self.async_db = self.async_client[db_name]
- self.async_collection = self.async_db[collection_name]
-
- def create_schema(self) -> None:
- """Create the database schema for the document manager."""
- pass
-
- async def acreate_schema(self) -> None:
- """Create the database schema for the document manager."""
- pass
-
- def update(
- self,
- keys: Sequence[str],
- *,
- group_ids: Optional[Sequence[Optional[str]]] = None,
- time_at_least: Optional[float] = None,
- ) -> None:
- """Upsert documents into the MongoDB collection."""
- if group_ids is None:
- group_ids = [None] * len(keys)
-
- if len(keys) != len(group_ids):
- raise ValueError("Number of keys does not match number of group_ids")
-
- for key, group_id in zip(keys, group_ids):
- self.sync_collection.find_one_and_update(
- {"namespace": self.namespace, "key": key},
- {"$set": {"group_id": group_id, "updated_at": self.get_time()}},
- upsert=True,
- )
-
- async def aupdate(
- self,
- keys: Sequence[str],
- *,
- group_ids: Optional[Sequence[Optional[str]]] = None,
- time_at_least: Optional[float] = None,
- ) -> None:
- """Asynchronously upsert documents into the MongoDB collection."""
- if group_ids is None:
- group_ids = [None] * len(keys)
-
- if len(keys) != len(group_ids):
- raise ValueError("Number of keys does not match number of group_ids")
-
- update_time = await self.aget_time()
- if time_at_least and update_time < time_at_least:
- raise ValueError("Server time is behind the expected time_at_least")
-
- for key, group_id in zip(keys, group_ids):
- await self.async_collection.find_one_and_update(
- {"namespace": self.namespace, "key": key},
- {"$set": {"group_id": group_id, "updated_at": update_time}},
- upsert=True,
- )
-
- def get_time(self) -> float:
- """Get the current server time as a timestamp."""
- server_info = self.sync_db.command("hostInfo")
- local_time = server_info["system"]["currentTime"]
- timestamp = local_time.timestamp()
- return timestamp
-
- async def aget_time(self) -> float:
- """Asynchronously get the current server time as a timestamp."""
- host_info = await self.async_collection.database.command("hostInfo")
- local_time = host_info["system"]["currentTime"]
- return local_time.timestamp()
-
- def exists(self, keys: Sequence[str]) -> List[bool]:
- """Check if the given keys exist in the MongoDB collection."""
- existing_keys = {
- doc["key"]
- for doc in self.sync_collection.find(
- {"namespace": self.namespace, "key": {"$in": keys}}, {"key": 1}
- )
- }
- return [key in existing_keys for key in keys]
-
- async def aexists(self, keys: Sequence[str]) -> List[bool]:
- """Asynchronously check if the given keys exist in the MongoDB collection."""
- cursor = self.async_collection.find(
- {"namespace": self.namespace, "key": {"$in": keys}}, {"key": 1}
- )
- existing_keys = {doc["key"] async for doc in cursor}
- return [key in existing_keys for key in keys]
-
- def list_keys(
- self,
- *,
- before: Optional[float] = None,
- after: Optional[float] = None,
- group_ids: Optional[Sequence[str]] = None,
- limit: Optional[int] = None,
- ) -> List[str]:
- """List documents in the MongoDB collection based on the provided date range."""
- query: Dict[str, Any] = {"namespace": self.namespace}
- if before:
- query["updated_at"] = {"$lt": before}
- if after:
- query["updated_at"] = {"$gt": after}
- if group_ids:
- query["group_id"] = {"$in": group_ids}
-
- cursor = (
- self.sync_collection.find(query, {"key": 1}).limit(limit)
- if limit
- else self.sync_collection.find(query, {"key": 1})
- )
- return [doc["key"] for doc in cursor]
-
- async def alist_keys(
- self,
- *,
- before: Optional[float] = None,
- after: Optional[float] = None,
- group_ids: Optional[Sequence[str]] = None,
- limit: Optional[int] = None,
- ) -> List[str]:
- """
- Asynchronously list documents in the MongoDB collection
- based on the provided date range.
- """
- query: Dict[str, Any] = {"namespace": self.namespace}
- if before:
- query["updated_at"] = {"$lt": before}
- if after:
- query["updated_at"] = {"$gt": after}
- if group_ids:
- query["group_id"] = {"$in": group_ids}
-
- cursor = (
- self.async_collection.find(query, {"key": 1}).limit(limit)
- if limit
- else self.async_collection.find(query, {"key": 1})
- )
- return [doc["key"] async for doc in cursor]
-
- def delete_keys(self, keys: Sequence[str]) -> None:
- """Delete documents from the MongoDB collection."""
- self.sync_collection.delete_many(
- {
- "namespace": self.namespace,
- "key": {"$in": keys},
- }
- )
-
- async def adelete_keys(self, keys: Sequence[str]) -> None:
- """Asynchronously delete documents from the MongoDB collection."""
- await self.async_collection.delete_many(
- {
- "namespace": self.namespace,
- "key": {"$in": keys},
- }
- )
diff --git a/libs/community/langchain_community/indexes/_sql_record_manager.py b/libs/community/langchain_community/indexes/_sql_record_manager.py
deleted file mode 100644
index bf59530a06..0000000000
--- a/libs/community/langchain_community/indexes/_sql_record_manager.py
+++ /dev/null
@@ -1,525 +0,0 @@
-"""Implementation of a record management layer in SQLAlchemy.
-
-The management layer uses SQLAlchemy to track upserted records.
-
-Currently, this layer only works with SQLite; hopwever, should be adaptable
-to other SQL implementations with minimal effort.
-
-Currently, includes an implementation that uses SQLAlchemy which should
-allow it to work with a variety of SQL as a backend.
-
-* Each key is associated with an updated_at field.
-* This filed is updated whenever the key is updated.
-* Keys can be listed based on the updated at field.
-* Keys can be deleted.
-"""
-
-import contextlib
-import decimal
-import uuid
-from typing import (
- Any,
- AsyncGenerator,
- Dict,
- Generator,
- List,
- Optional,
- Sequence,
- Union,
- cast,
-)
-
-from sqlalchemy import (
- Column,
- Float,
- Index,
- String,
- UniqueConstraint,
- and_,
- create_engine,
- delete,
- select,
- text,
-)
-from sqlalchemy.engine import URL, Engine
-from sqlalchemy.ext.asyncio import (
- AsyncEngine,
- AsyncSession,
- create_async_engine,
-)
-from sqlalchemy.ext.declarative import declarative_base
-from sqlalchemy.orm import Session, sessionmaker
-
-try:
- from sqlalchemy.ext.asyncio import async_sessionmaker
-except ImportError:
- # dummy for sqlalchemy < 2
- async_sessionmaker = type("async_sessionmaker", (type,), {}) # type: ignore[assignment,misc]
-
-from langchain_community.indexes.base import RecordManager
-
-Base = declarative_base()
-
-
-class UpsertionRecord(Base): # type: ignore[valid-type,misc]
- """Table used to keep track of when a key was last updated."""
-
- # ATTENTION:
- # Prior to modifying this table, please determine whether
- # we should create migrations for this table to make sure
- # users do not experience data loss.
- __tablename__ = "upsertion_record"
-
- uuid = Column(
- String,
- index=True,
- default=lambda: str(uuid.uuid4()),
- primary_key=True,
- nullable=False,
- )
- key = Column(String, index=True)
- # Using a non-normalized representation to handle `namespace` attribute.
- # If the need arises, this attribute can be pulled into a separate Collection
- # table at some time later.
- namespace = Column(String, index=True, nullable=False)
- group_id = Column(String, index=True, nullable=True)
-
- # The timestamp associated with the last record upsertion.
- updated_at = Column(Float, index=True)
-
- __table_args__ = (
- UniqueConstraint("key", "namespace", name="uix_key_namespace"),
- Index("ix_key_namespace", "key", "namespace"),
- )
-
-
-class SQLRecordManager(RecordManager):
- """A SQL Alchemy based implementation of the record manager."""
-
- def __init__(
- self,
- namespace: str,
- *,
- engine: Optional[Union[Engine, AsyncEngine]] = None,
- db_url: Union[None, str, URL] = None,
- engine_kwargs: Optional[Dict[str, Any]] = None,
- async_mode: bool = False,
- ) -> None:
- """Initialize the SQLRecordManager.
-
- This class serves as a manager persistence layer that uses an SQL
- backend to track upserted records. You should specify either a db_url
- to create an engine or provide an existing engine.
-
- Args:
- namespace: The namespace associated with this record manager.
- engine: An already existing SQL Alchemy engine.
- Default is None.
- db_url: A database connection string used to create
- an SQL Alchemy engine. Default is None.
- engine_kwargs: Additional keyword arguments
- to be passed when creating the engine. Default is an empty dictionary.
- async_mode: Whether to create an async engine.
- Driver should support async operations.
- It only applies if db_url is provided.
- Default is False.
-
- Raises:
- ValueError: If both db_url and engine are provided or neither.
- AssertionError: If something unexpected happens during engine configuration.
- """
- super().__init__(namespace=namespace)
- if db_url is None and engine is None:
- raise ValueError("Must specify either db_url or engine")
-
- if db_url is not None and engine is not None:
- raise ValueError("Must specify either db_url or engine, not both")
-
- _engine: Union[Engine, AsyncEngine]
- if db_url:
- if async_mode:
- _engine = create_async_engine(db_url, **(engine_kwargs or {}))
- else:
- _engine = create_engine(db_url, **(engine_kwargs or {}))
- elif engine:
- _engine = engine
-
- else:
- raise AssertionError("Something went wrong with configuration of engine.")
-
- _session_factory: Union[sessionmaker[Session], async_sessionmaker[AsyncSession]]
- if isinstance(_engine, AsyncEngine):
- _session_factory = async_sessionmaker(bind=_engine)
- else:
- _session_factory = sessionmaker(bind=_engine)
-
- self.engine = _engine
- self.dialect = _engine.dialect.name
- self.session_factory = _session_factory
-
- def create_schema(self) -> None:
- """Create the database schema."""
- if isinstance(self.engine, AsyncEngine):
- raise AssertionError("This method is not supported for async engines.")
-
- Base.metadata.create_all(self.engine)
-
- async def acreate_schema(self) -> None:
- """Create the database schema."""
-
- if not isinstance(self.engine, AsyncEngine):
- raise AssertionError("This method is not supported for sync engines.")
-
- async with self.engine.begin() as session:
- await session.run_sync(Base.metadata.create_all)
-
- @contextlib.contextmanager
- def _make_session(self) -> Generator[Session, None, None]:
- """Create a session and close it after use."""
-
- if isinstance(self.session_factory, async_sessionmaker):
- raise AssertionError("This method is not supported for async engines.")
-
- session = self.session_factory()
- try:
- yield session
- finally:
- session.close()
-
- @contextlib.asynccontextmanager
- async def _amake_session(self) -> AsyncGenerator[AsyncSession, None]:
- """Create a session and close it after use."""
-
- if not isinstance(self.engine, AsyncEngine):
- raise AssertionError("This method is not supported for sync engines.")
-
- async with cast(AsyncSession, self.session_factory()) as session:
- yield session
-
- def get_time(self) -> float:
- """Get the current server time as a timestamp.
-
- Please note it's critical that time is obtained from the server since
- we want a monotonic clock.
- """
- with self._make_session() as session:
- # * SQLite specific implementation, can be changed based on dialect.
- # * For SQLite, unlike unixepoch it will work with older versions of SQLite.
- # ----
- # julianday('now'): Julian day number for the current date and time.
- # The Julian day is a continuous count of days, starting from a
- # reference date (Julian day number 0).
- # 2440587.5 - constant represents the Julian day number for January 1, 1970
- # 86400.0 - constant represents the number of seconds
- # in a day (24 hours * 60 minutes * 60 seconds)
- if self.dialect == "sqlite":
- query = text("SELECT (julianday('now') - 2440587.5) * 86400.0;")
- elif self.dialect == "postgresql":
- query = text("SELECT EXTRACT (EPOCH FROM CURRENT_TIMESTAMP);")
- else:
- raise NotImplementedError(f"Not implemented for dialect {self.dialect}")
-
- dt = session.execute(query).scalar()
- if isinstance(dt, decimal.Decimal):
- dt = float(dt)
- if not isinstance(dt, float):
- raise AssertionError(f"Unexpected type for datetime: {type(dt)}")
- return dt
-
- async def aget_time(self) -> float:
- """Get the current server time as a timestamp.
-
- Please note it's critical that time is obtained from the server since
- we want a monotonic clock.
- """
- async with self._amake_session() as session:
- # * SQLite specific implementation, can be changed based on dialect.
- # * For SQLite, unlike unixepoch it will work with older versions of SQLite.
- # ----
- # julianday('now'): Julian day number for the current date and time.
- # The Julian day is a continuous count of days, starting from a
- # reference date (Julian day number 0).
- # 2440587.5 - constant represents the Julian day number for January 1, 1970
- # 86400.0 - constant represents the number of seconds
- # in a day (24 hours * 60 minutes * 60 seconds)
- if self.dialect == "sqlite":
- query = text("SELECT (julianday('now') - 2440587.5) * 86400.0;")
- elif self.dialect == "postgresql":
- query = text("SELECT EXTRACT (EPOCH FROM CURRENT_TIMESTAMP);")
- else:
- raise NotImplementedError(f"Not implemented for dialect {self.dialect}")
-
- dt = (await session.execute(query)).scalar_one_or_none()
-
- if isinstance(dt, decimal.Decimal):
- dt = float(dt)
- if not isinstance(dt, float):
- raise AssertionError(f"Unexpected type for datetime: {type(dt)}")
- return dt
-
- def update(
- self,
- keys: Sequence[str],
- *,
- group_ids: Optional[Sequence[Optional[str]]] = None,
- time_at_least: Optional[float] = None,
- ) -> None:
- """Upsert records into the SQLite database."""
- if group_ids is None:
- group_ids = [None] * len(keys)
-
- if len(keys) != len(group_ids):
- raise ValueError(
- f"Number of keys ({len(keys)}) does not match number of "
- f"group_ids ({len(group_ids)})"
- )
-
- # Get the current time from the server.
- # This makes an extra round trip to the server, should not be a big deal
- # if the batch size is large enough.
- # Getting the time here helps us compare it against the time_at_least
- # and raise an error if there is a time sync issue.
- # Here, we're just being extra careful to minimize the chance of
- # data loss due to incorrectly deleting records.
- update_time = self.get_time()
-
- if time_at_least and update_time < time_at_least:
- # Safeguard against time sync issues
- raise AssertionError(f"Time sync issue: {update_time} < {time_at_least}")
-
- records_to_upsert = [
- {
- "key": key,
- "namespace": self.namespace,
- "updated_at": update_time,
- "group_id": group_id,
- }
- for key, group_id in zip(keys, group_ids)
- ]
-
- with self._make_session() as session:
- if self.dialect == "sqlite":
- from sqlalchemy.dialects.sqlite import insert as sqlite_insert
-
- # Note: uses SQLite insert to make on_conflict_do_update work.
- # This code needs to be generalized a bit to work with more dialects.
- insert_stmt = sqlite_insert(UpsertionRecord).values(records_to_upsert)
- stmt = insert_stmt.on_conflict_do_update(
- [UpsertionRecord.key, UpsertionRecord.namespace],
- set_=dict(
- # attr-defined type ignore
- updated_at=insert_stmt.excluded.updated_at,
- group_id=insert_stmt.excluded.group_id,
- ),
- )
- elif self.dialect == "postgresql":
- from sqlalchemy.dialects.postgresql import insert as pg_insert
-
- # Note: uses SQLite insert to make on_conflict_do_update work.
- # This code needs to be generalized a bit to work with more dialects.
- insert_stmt = pg_insert(UpsertionRecord).values(records_to_upsert) # type: ignore[assignment]
- stmt = insert_stmt.on_conflict_do_update(
- "uix_key_namespace", # Name of constraint
- set_=dict(
- # attr-defined type ignore
- updated_at=insert_stmt.excluded.updated_at,
- group_id=insert_stmt.excluded.group_id,
- ),
- )
- else:
- raise NotImplementedError(f"Unsupported dialect {self.dialect}")
-
- session.execute(stmt)
- session.commit()
-
- async def aupdate(
- self,
- keys: Sequence[str],
- *,
- group_ids: Optional[Sequence[Optional[str]]] = None,
- time_at_least: Optional[float] = None,
- ) -> None:
- """Upsert records into the SQLite database."""
- if group_ids is None:
- group_ids = [None] * len(keys)
-
- if len(keys) != len(group_ids):
- raise ValueError(
- f"Number of keys ({len(keys)}) does not match number of "
- f"group_ids ({len(group_ids)})"
- )
-
- # Get the current time from the server.
- # This makes an extra round trip to the server, should not be a big deal
- # if the batch size is large enough.
- # Getting the time here helps us compare it against the time_at_least
- # and raise an error if there is a time sync issue.
- # Here, we're just being extra careful to minimize the chance of
- # data loss due to incorrectly deleting records.
- update_time = await self.aget_time()
-
- if time_at_least and update_time < time_at_least:
- # Safeguard against time sync issues
- raise AssertionError(f"Time sync issue: {update_time} < {time_at_least}")
-
- records_to_upsert = [
- {
- "key": key,
- "namespace": self.namespace,
- "updated_at": update_time,
- "group_id": group_id,
- }
- for key, group_id in zip(keys, group_ids)
- ]
-
- async with self._amake_session() as session:
- if self.dialect == "sqlite":
- from sqlalchemy.dialects.sqlite import insert as sqlite_insert
-
- # Note: uses SQLite insert to make on_conflict_do_update work.
- # This code needs to be generalized a bit to work with more dialects.
- insert_stmt = sqlite_insert(UpsertionRecord).values(records_to_upsert)
- stmt = insert_stmt.on_conflict_do_update(
- [UpsertionRecord.key, UpsertionRecord.namespace],
- set_=dict(
- # attr-defined type ignore
- updated_at=insert_stmt.excluded.updated_at,
- group_id=insert_stmt.excluded.group_id,
- ),
- )
- elif self.dialect == "postgresql":
- from sqlalchemy.dialects.postgresql import insert as pg_insert
-
- # Note: uses SQLite insert to make on_conflict_do_update work.
- # This code needs to be generalized a bit to work with more dialects.
- insert_stmt = pg_insert(UpsertionRecord).values(records_to_upsert) # type: ignore[assignment]
- stmt = insert_stmt.on_conflict_do_update(
- "uix_key_namespace", # Name of constraint
- set_=dict(
- # attr-defined type ignore
- updated_at=insert_stmt.excluded.updated_at,
- group_id=insert_stmt.excluded.group_id,
- ),
- )
- else:
- raise NotImplementedError(f"Unsupported dialect {self.dialect}")
-
- await session.execute(stmt)
- await session.commit()
-
- def exists(self, keys: Sequence[str]) -> List[bool]:
- """Check if the given keys exist in the SQLite database."""
- with self._make_session() as session:
- records = (
- # mypy does not recognize .all()
- session.query(UpsertionRecord.key)
- .filter(
- and_(
- UpsertionRecord.key.in_(keys),
- UpsertionRecord.namespace == self.namespace,
- )
- )
- .all()
- )
- found_keys = set(r.key for r in records)
- return [k in found_keys for k in keys]
-
- async def aexists(self, keys: Sequence[str]) -> List[bool]:
- """Check if the given keys exist in the SQLite database."""
- async with self._amake_session() as session:
- records = (
- (
- await session.execute(
- select(UpsertionRecord.key).where(
- and_(
- UpsertionRecord.key.in_(keys),
- UpsertionRecord.namespace == self.namespace,
- )
- )
- )
- )
- .scalars()
- .all()
- )
- found_keys = set(records)
- return [k in found_keys for k in keys]
-
- def list_keys(
- self,
- *,
- before: Optional[float] = None,
- after: Optional[float] = None,
- group_ids: Optional[Sequence[str]] = None,
- limit: Optional[int] = None,
- ) -> List[str]:
- """List records in the SQLite database based on the provided date range."""
- with self._make_session() as session:
- query = session.query(UpsertionRecord).filter(
- UpsertionRecord.namespace == self.namespace
- )
-
- # mypy does not recognize .all() or .filter()
- if after:
- query = query.filter(UpsertionRecord.updated_at > after)
- if before:
- query = query.filter(UpsertionRecord.updated_at < before)
- if group_ids:
- query = query.filter(UpsertionRecord.group_id.in_(group_ids))
-
- if limit:
- query = query.limit(limit)
- records = query.all()
- return [r.key for r in records] # type: ignore[misc]
-
- async def alist_keys(
- self,
- *,
- before: Optional[float] = None,
- after: Optional[float] = None,
- group_ids: Optional[Sequence[str]] = None,
- limit: Optional[int] = None,
- ) -> List[str]:
- """List records in the SQLite database based on the provided date range."""
- async with self._amake_session() as session:
- query = select(UpsertionRecord.key).filter(
- UpsertionRecord.namespace == self.namespace
- )
-
- # mypy does not recognize .all() or .filter()
- if after:
- query = query.filter(UpsertionRecord.updated_at > after)
- if before:
- query = query.filter(UpsertionRecord.updated_at < before)
- if group_ids:
- query = query.filter(UpsertionRecord.group_id.in_(group_ids))
-
- if limit:
- query = query.limit(limit)
- records = (await session.execute(query)).scalars().all()
- return list(records)
-
- def delete_keys(self, keys: Sequence[str]) -> None:
- """Delete records from the SQLite database."""
- with self._make_session() as session:
- # mypy does not recognize .delete()
- session.query(UpsertionRecord).filter(
- and_(
- UpsertionRecord.key.in_(keys),
- UpsertionRecord.namespace == self.namespace,
- )
- ).delete()
- session.commit()
-
- async def adelete_keys(self, keys: Sequence[str]) -> None:
- """Delete records from the SQLite database."""
- async with self._amake_session() as session:
- await session.execute(
- delete(UpsertionRecord).where(
- and_(
- UpsertionRecord.key.in_(keys),
- UpsertionRecord.namespace == self.namespace,
- )
- )
- )
-
- await session.commit()
diff --git a/libs/community/langchain_community/indexes/base.py b/libs/community/langchain_community/indexes/base.py
deleted file mode 100644
index 97805d91e7..0000000000
--- a/libs/community/langchain_community/indexes/base.py
+++ /dev/null
@@ -1,172 +0,0 @@
-from __future__ import annotations
-
-import uuid
-from abc import ABC, abstractmethod
-from typing import List, Optional, Sequence
-
-NAMESPACE_UUID = uuid.UUID(int=1984)
-
-
-class RecordManager(ABC):
- """Abstract base class for a record manager."""
-
- def __init__(
- self,
- namespace: str,
- ) -> None:
- """Initialize the record manager.
-
- Args:
- namespace (str): The namespace for the record manager.
- """
- self.namespace = namespace
-
- @abstractmethod
- def create_schema(self) -> None:
- """Create the database schema for the record manager."""
-
- @abstractmethod
- async def acreate_schema(self) -> None:
- """Create the database schema for the record manager."""
-
- @abstractmethod
- def get_time(self) -> float:
- """Get the current server time as a high resolution timestamp!
-
- It's important to get this from the server to ensure a monotonic clock,
- otherwise there may be data loss when cleaning up old documents!
-
- Returns:
- The current server time as a float timestamp.
- """
-
- @abstractmethod
- async def aget_time(self) -> float:
- """Get the current server time as a high resolution timestamp!
-
- It's important to get this from the server to ensure a monotonic clock,
- otherwise there may be data loss when cleaning up old documents!
-
- Returns:
- The current server time as a float timestamp.
- """
-
- @abstractmethod
- def update(
- self,
- keys: Sequence[str],
- *,
- group_ids: Optional[Sequence[Optional[str]]] = None,
- time_at_least: Optional[float] = None,
- ) -> None:
- """Upsert records into the database.
-
- Args:
- keys: A list of record keys to upsert.
- group_ids: A list of group IDs corresponding to the keys.
- time_at_least: if provided, updates should only happen if the
- updated_at field is at least this time.
-
- Raises:
- ValueError: If the length of keys doesn't match the length of group_ids.
- """
-
- @abstractmethod
- async def aupdate(
- self,
- keys: Sequence[str],
- *,
- group_ids: Optional[Sequence[Optional[str]]] = None,
- time_at_least: Optional[float] = None,
- ) -> None:
- """Upsert records into the database.
-
- Args:
- keys: A list of record keys to upsert.
- group_ids: A list of group IDs corresponding to the keys.
- time_at_least: if provided, updates should only happen if the
- updated_at field is at least this time.
-
- Raises:
- ValueError: If the length of keys doesn't match the length of group_ids.
- """
-
- @abstractmethod
- def exists(self, keys: Sequence[str]) -> List[bool]:
- """Check if the provided keys exist in the database.
-
- Args:
- keys: A list of keys to check.
-
- Returns:
- A list of boolean values indicating the existence of each key.
- """
-
- @abstractmethod
- async def aexists(self, keys: Sequence[str]) -> List[bool]:
- """Check if the provided keys exist in the database.
-
- Args:
- keys: A list of keys to check.
-
- Returns:
- A list of boolean values indicating the existence of each key.
- """
-
- @abstractmethod
- def list_keys(
- self,
- *,
- before: Optional[float] = None,
- after: Optional[float] = None,
- group_ids: Optional[Sequence[str]] = None,
- limit: Optional[int] = None,
- ) -> List[str]:
- """List records in the database based on the provided filters.
-
- Args:
- before: Filter to list records updated before this time.
- after: Filter to list records updated after this time.
- group_ids: Filter to list records with specific group IDs.
- limit: optional limit on the number of records to return.
-
- Returns:
- A list of keys for the matching records.
- """
-
- @abstractmethod
- async def alist_keys(
- self,
- *,
- before: Optional[float] = None,
- after: Optional[float] = None,
- group_ids: Optional[Sequence[str]] = None,
- limit: Optional[int] = None,
- ) -> List[str]:
- """List records in the database based on the provided filters.
-
- Args:
- before: Filter to list records updated before this time.
- after: Filter to list records updated after this time.
- group_ids: Filter to list records with specific group IDs.
- limit: optional limit on the number of records to return.
-
- Returns:
- A list of keys for the matching records.
- """
-
- @abstractmethod
- def delete_keys(self, keys: Sequence[str]) -> None:
- """Delete specified records from the database.
-
- Args:
- keys: A list of keys to delete.
- """
-
- @abstractmethod
- async def adelete_keys(self, keys: Sequence[str]) -> None:
- """Delete specified records from the database.
-
- Args:
- keys: A list of keys to delete.
- """
diff --git a/libs/community/langchain_community/llms/__init__.py b/libs/community/langchain_community/llms/__init__.py
deleted file mode 100644
index 45e0052429..0000000000
--- a/libs/community/langchain_community/llms/__init__.py
+++ /dev/null
@@ -1,1102 +0,0 @@
-"""
-**LLM** classes provide
-access to the large language model (**LLM**) APIs and services.
-
-**Class hierarchy:**
-
-.. code-block::
-
- BaseLanguageModel --> BaseLLM --> LLM --> # Examples: AI21, HuggingFaceHub, OpenAI
-
-**Main helpers:**
-
-.. code-block::
-
- LLMResult, PromptValue,
- CallbackManagerForLLMRun, AsyncCallbackManagerForLLMRun,
- CallbackManager, AsyncCallbackManager,
- AIMessage, BaseMessage
-""" # noqa: E501
-
-from typing import Any, Callable, Dict, Type
-
-from langchain_core._api.deprecation import warn_deprecated
-from langchain_core.language_models.llms import BaseLLM
-
-
-def _import_ai21() -> Type[BaseLLM]:
- from langchain_community.llms.ai21 import AI21
-
- return AI21
-
-
-def _import_aleph_alpha() -> Type[BaseLLM]:
- from langchain_community.llms.aleph_alpha import AlephAlpha
-
- return AlephAlpha
-
-
-def _import_amazon_api_gateway() -> Type[BaseLLM]:
- from langchain_community.llms.amazon_api_gateway import AmazonAPIGateway
-
- return AmazonAPIGateway
-
-
-def _import_anthropic() -> Type[BaseLLM]:
- from langchain_community.llms.anthropic import Anthropic
-
- return Anthropic
-
-
-def _import_anyscale() -> Type[BaseLLM]:
- from langchain_community.llms.anyscale import Anyscale
-
- return Anyscale
-
-
-def _import_aphrodite() -> Type[BaseLLM]:
- from langchain_community.llms.aphrodite import Aphrodite
-
- return Aphrodite
-
-
-def _import_arcee() -> Type[BaseLLM]:
- from langchain_community.llms.arcee import Arcee
-
- return Arcee
-
-
-def _import_aviary() -> Type[BaseLLM]:
- from langchain_community.llms.aviary import Aviary
-
- return Aviary
-
-
-def _import_azureml_endpoint() -> Type[BaseLLM]:
- from langchain_community.llms.azureml_endpoint import AzureMLOnlineEndpoint
-
- return AzureMLOnlineEndpoint
-
-
-def _import_baichuan() -> Type[BaseLLM]:
- from langchain_community.llms.baichuan import BaichuanLLM
-
- return BaichuanLLM
-
-
-def _import_baidu_qianfan_endpoint() -> Type[BaseLLM]:
- from langchain_community.llms.baidu_qianfan_endpoint import QianfanLLMEndpoint
-
- return QianfanLLMEndpoint
-
-
-def _import_bananadev() -> Type[BaseLLM]:
- from langchain_community.llms.bananadev import Banana
-
- return Banana
-
-
-def _import_baseten() -> Type[BaseLLM]:
- from langchain_community.llms.baseten import Baseten
-
- return Baseten
-
-
-def _import_beam() -> Type[BaseLLM]:
- from langchain_community.llms.beam import Beam
-
- return Beam
-
-
-def _import_bedrock() -> Type[BaseLLM]:
- from langchain_community.llms.bedrock import Bedrock
-
- return Bedrock
-
-
-def _import_bigdlllm() -> Type[BaseLLM]:
- from langchain_community.llms.bigdl_llm import BigdlLLM
-
- return BigdlLLM
-
-
-def _import_bittensor() -> Type[BaseLLM]:
- from langchain_community.llms.bittensor import NIBittensorLLM
-
- return NIBittensorLLM
-
-
-def _import_cerebriumai() -> Type[BaseLLM]:
- from langchain_community.llms.cerebriumai import CerebriumAI
-
- return CerebriumAI
-
-
-def _import_chatglm() -> Type[BaseLLM]:
- from langchain_community.llms.chatglm import ChatGLM
-
- return ChatGLM
-
-
-def _import_clarifai() -> Type[BaseLLM]:
- from langchain_community.llms.clarifai import Clarifai
-
- return Clarifai
-
-
-def _import_cohere() -> Type[BaseLLM]:
- from langchain_community.llms.cohere import Cohere
-
- return Cohere
-
-
-def _import_ctransformers() -> Type[BaseLLM]:
- from langchain_community.llms.ctransformers import CTransformers
-
- return CTransformers
-
-
-def _import_ctranslate2() -> Type[BaseLLM]:
- from langchain_community.llms.ctranslate2 import CTranslate2
-
- return CTranslate2
-
-
-def _import_databricks() -> Type[BaseLLM]:
- from langchain_community.llms.databricks import Databricks
-
- return Databricks
-
-
-# deprecated / only for back compat - do not add to __all__
-def _import_databricks_chat() -> Any:
- warn_deprecated(
- since="0.0.22",
- removal="1.0",
- alternative_import="langchain_community.chat_models.ChatDatabricks",
- )
- from langchain_community.chat_models.databricks import ChatDatabricks
-
- return ChatDatabricks
-
-
-def _import_deepinfra() -> Type[BaseLLM]:
- from langchain_community.llms.deepinfra import DeepInfra
-
- return DeepInfra
-
-
-def _import_deepsparse() -> Type[BaseLLM]:
- from langchain_community.llms.deepsparse import DeepSparse
-
- return DeepSparse
-
-
-def _import_edenai() -> Type[BaseLLM]:
- from langchain_community.llms.edenai import EdenAI
-
- return EdenAI
-
-
-def _import_fake() -> Type[BaseLLM]:
- from langchain_community.llms.fake import FakeListLLM
-
- return FakeListLLM
-
-
-def _import_fireworks() -> Type[BaseLLM]:
- from langchain_community.llms.fireworks import Fireworks
-
- return Fireworks
-
-
-def _import_forefrontai() -> Type[BaseLLM]:
- from langchain_community.llms.forefrontai import ForefrontAI
-
- return ForefrontAI
-
-
-def _import_friendli() -> Type[BaseLLM]:
- from langchain_community.llms.friendli import Friendli
-
- return Friendli
-
-
-def _import_gigachat() -> Type[BaseLLM]:
- from langchain_community.llms.gigachat import GigaChat
-
- return GigaChat
-
-
-def _import_google_palm() -> Type[BaseLLM]:
- from langchain_community.llms.google_palm import GooglePalm
-
- return GooglePalm
-
-
-def _import_gooseai() -> Type[BaseLLM]:
- from langchain_community.llms.gooseai import GooseAI
-
- return GooseAI
-
-
-def _import_gpt4all() -> Type[BaseLLM]:
- from langchain_community.llms.gpt4all import GPT4All
-
- return GPT4All
-
-
-def _import_gradient_ai() -> Type[BaseLLM]:
- from langchain_community.llms.gradient_ai import GradientLLM
-
- return GradientLLM
-
-
-def _import_huggingface_endpoint() -> Type[BaseLLM]:
- from langchain_community.llms.huggingface_endpoint import HuggingFaceEndpoint
-
- return HuggingFaceEndpoint
-
-
-def _import_huggingface_hub() -> Type[BaseLLM]:
- from langchain_community.llms.huggingface_hub import HuggingFaceHub
-
- return HuggingFaceHub
-
-
-def _import_huggingface_pipeline() -> Type[BaseLLM]:
- from langchain_community.llms.huggingface_pipeline import HuggingFacePipeline
-
- return HuggingFacePipeline
-
-
-def _import_huggingface_text_gen_inference() -> Type[BaseLLM]:
- from langchain_community.llms.huggingface_text_gen_inference import (
- HuggingFaceTextGenInference,
- )
-
- return HuggingFaceTextGenInference
-
-
-def _import_human() -> Type[BaseLLM]:
- from langchain_community.llms.human import HumanInputLLM
-
- return HumanInputLLM
-
-
-def _import_ipex_llm() -> Type[BaseLLM]:
- from langchain_community.llms.ipex_llm import IpexLLM
-
- return IpexLLM
-
-
-def _import_javelin_ai_gateway() -> Type[BaseLLM]:
- from langchain_community.llms.javelin_ai_gateway import JavelinAIGateway
-
- return JavelinAIGateway
-
-
-def _import_koboldai() -> Type[BaseLLM]:
- from langchain_community.llms.koboldai import KoboldApiLLM
-
- return KoboldApiLLM
-
-
-def _import_konko() -> Type[BaseLLM]:
- from langchain_community.llms.konko import Konko
-
- return Konko
-
-
-def _import_llamacpp() -> Type[BaseLLM]:
- from langchain_community.llms.llamacpp import LlamaCpp
-
- return LlamaCpp
-
-
-def _import_llamafile() -> Type[BaseLLM]:
- from langchain_community.llms.llamafile import Llamafile
-
- return Llamafile
-
-
-def _import_manifest() -> Type[BaseLLM]:
- from langchain_community.llms.manifest import ManifestWrapper
-
- return ManifestWrapper
-
-
-def _import_minimax() -> Type[BaseLLM]:
- from langchain_community.llms.minimax import Minimax
-
- return Minimax
-
-
-def _import_mlflow() -> Type[BaseLLM]:
- from langchain_community.llms.mlflow import Mlflow
-
- return Mlflow
-
-
-# deprecated / only for back compat - do not add to __all__
-def _import_mlflow_chat() -> Any:
- warn_deprecated(
- since="0.0.22",
- removal="1.0",
- alternative_import="langchain_community.chat_models.ChatMlflow",
- )
- from langchain_community.chat_models.mlflow import ChatMlflow
-
- return ChatMlflow
-
-
-def _import_mlflow_ai_gateway() -> Type[BaseLLM]:
- from langchain_community.llms.mlflow_ai_gateway import MlflowAIGateway
-
- return MlflowAIGateway
-
-
-def _import_mlx_pipeline() -> Type[BaseLLM]:
- from langchain_community.llms.mlx_pipeline import MLXPipeline
-
- return MLXPipeline
-
-
-def _import_modal() -> Type[BaseLLM]:
- from langchain_community.llms.modal import Modal
-
- return Modal
-
-
-def _import_mosaicml() -> Type[BaseLLM]:
- from langchain_community.llms.mosaicml import MosaicML
-
- return MosaicML
-
-
-def _import_nlpcloud() -> Type[BaseLLM]:
- from langchain_community.llms.nlpcloud import NLPCloud
-
- return NLPCloud
-
-
-def _import_oci_md_tgi() -> Type[BaseLLM]:
- from langchain_community.llms.oci_data_science_model_deployment_endpoint import (
- OCIModelDeploymentTGI,
- )
-
- return OCIModelDeploymentTGI
-
-
-def _import_oci_md_vllm() -> Type[BaseLLM]:
- from langchain_community.llms.oci_data_science_model_deployment_endpoint import (
- OCIModelDeploymentVLLM,
- )
-
- return OCIModelDeploymentVLLM
-
-
-def _import_oci_md() -> Type[BaseLLM]:
- from langchain_community.llms.oci_data_science_model_deployment_endpoint import (
- OCIModelDeploymentLLM,
- )
-
- return OCIModelDeploymentLLM
-
-
-def _import_oci_gen_ai() -> Type[BaseLLM]:
- from langchain_community.llms.oci_generative_ai import OCIGenAI
-
- return OCIGenAI
-
-
-def _import_octoai_endpoint() -> Type[BaseLLM]:
- from langchain_community.llms.octoai_endpoint import OctoAIEndpoint
-
- return OctoAIEndpoint
-
-
-def _import_ollama() -> Type[BaseLLM]:
- from langchain_community.llms.ollama import Ollama
-
- return Ollama
-
-
-def _import_opaqueprompts() -> Type[BaseLLM]:
- from langchain_community.llms.opaqueprompts import OpaquePrompts
-
- return OpaquePrompts
-
-
-def _import_azure_openai() -> Type[BaseLLM]:
- from langchain_community.llms.openai import AzureOpenAI
-
- return AzureOpenAI
-
-
-def _import_openai() -> Type[BaseLLM]:
- from langchain_community.llms.openai import OpenAI
-
- return OpenAI
-
-
-def _import_openai_chat() -> Type[BaseLLM]:
- from langchain_community.llms.openai import OpenAIChat
-
- return OpenAIChat
-
-
-def _import_openllm() -> Type[BaseLLM]:
- from langchain_community.llms.openllm import OpenLLM
-
- return OpenLLM
-
-
-def _import_openlm() -> Type[BaseLLM]:
- from langchain_community.llms.openlm import OpenLM
-
- return OpenLM
-
-
-def _import_outlines() -> Type[BaseLLM]:
- from langchain_community.llms.outlines import Outlines
-
- return Outlines
-
-
-def _import_pai_eas_endpoint() -> Type[BaseLLM]:
- from langchain_community.llms.pai_eas_endpoint import PaiEasEndpoint
-
- return PaiEasEndpoint
-
-
-def _import_petals() -> Type[BaseLLM]:
- from langchain_community.llms.petals import Petals
-
- return Petals
-
-
-def _import_pipelineai() -> Type[BaseLLM]:
- from langchain_community.llms.pipelineai import PipelineAI
-
- return PipelineAI
-
-
-def _import_predibase() -> Type[BaseLLM]:
- from langchain_community.llms.predibase import Predibase
-
- return Predibase
-
-
-def _import_predictionguard() -> Type[BaseLLM]:
- from langchain_community.llms.predictionguard import PredictionGuard
-
- return PredictionGuard
-
-
-def _import_promptlayer() -> Type[BaseLLM]:
- from langchain_community.llms.promptlayer_openai import PromptLayerOpenAI
-
- return PromptLayerOpenAI
-
-
-def _import_promptlayer_chat() -> Type[BaseLLM]:
- from langchain_community.llms.promptlayer_openai import PromptLayerOpenAIChat
-
- return PromptLayerOpenAIChat
-
-
-def _import_replicate() -> Type[BaseLLM]:
- from langchain_community.llms.replicate import Replicate
-
- return Replicate
-
-
-def _import_rwkv() -> Type[BaseLLM]:
- from langchain_community.llms.rwkv import RWKV
-
- return RWKV
-
-
-def _import_sagemaker_endpoint() -> Type[BaseLLM]:
- from langchain_community.llms.sagemaker_endpoint import SagemakerEndpoint
-
- return SagemakerEndpoint
-
-
-def _import_sambanovacloud() -> Type[BaseLLM]:
- from langchain_community.llms.sambanova import SambaNovaCloud
-
- return SambaNovaCloud
-
-
-def _import_sambastudio() -> Type[BaseLLM]:
- from langchain_community.llms.sambanova import SambaStudio
-
- return SambaStudio
-
-
-def _import_self_hosted() -> Type[BaseLLM]:
- from langchain_community.llms.self_hosted import SelfHostedPipeline
-
- return SelfHostedPipeline
-
-
-def _import_self_hosted_hugging_face() -> Type[BaseLLM]:
- from langchain_community.llms.self_hosted_hugging_face import (
- SelfHostedHuggingFaceLLM,
- )
-
- return SelfHostedHuggingFaceLLM
-
-
-def _import_stochasticai() -> Type[BaseLLM]:
- from langchain_community.llms.stochasticai import StochasticAI
-
- return StochasticAI
-
-
-def _import_symblai_nebula() -> Type[BaseLLM]:
- from langchain_community.llms.symblai_nebula import Nebula
-
- return Nebula
-
-
-def _import_textgen() -> Type[BaseLLM]:
- from langchain_community.llms.textgen import TextGen
-
- return TextGen
-
-
-def _import_titan_takeoff() -> Type[BaseLLM]:
- from langchain_community.llms.titan_takeoff import TitanTakeoff
-
- return TitanTakeoff
-
-
-def _import_titan_takeoff_pro() -> Type[BaseLLM]:
- from langchain_community.llms.titan_takeoff import TitanTakeoff
-
- return TitanTakeoff
-
-
-def _import_together() -> Type[BaseLLM]:
- from langchain_community.llms.together import Together
-
- return Together
-
-
-def _import_tongyi() -> Type[BaseLLM]:
- from langchain_community.llms.tongyi import Tongyi
-
- return Tongyi
-
-
-def _import_vertex() -> Type[BaseLLM]:
- from langchain_community.llms.vertexai import VertexAI
-
- return VertexAI
-
-
-def _import_vertex_model_garden() -> Type[BaseLLM]:
- from langchain_community.llms.vertexai import VertexAIModelGarden
-
- return VertexAIModelGarden
-
-
-def _import_vllm() -> Type[BaseLLM]:
- from langchain_community.llms.vllm import VLLM
-
- return VLLM
-
-
-def _import_vllm_openai() -> Type[BaseLLM]:
- from langchain_community.llms.vllm import VLLMOpenAI
-
- return VLLMOpenAI
-
-
-def _import_watsonxllm() -> Type[BaseLLM]:
- from langchain_community.llms.watsonxllm import WatsonxLLM
-
- return WatsonxLLM
-
-
-def _import_weight_only_quantization() -> Any:
- from langchain_community.llms.weight_only_quantization import (
- WeightOnlyQuantPipeline,
- )
-
- return WeightOnlyQuantPipeline
-
-
-def _import_writer() -> Type[BaseLLM]:
- from langchain_community.llms.writer import Writer
-
- return Writer
-
-
-def _import_xinference() -> Type[BaseLLM]:
- from langchain_community.llms.xinference import Xinference
-
- return Xinference
-
-
-def _import_yandex_gpt() -> Type[BaseLLM]:
- from langchain_community.llms.yandex import YandexGPT
-
- return YandexGPT
-
-
-def _import_yuan2() -> Type[BaseLLM]:
- from langchain_community.llms.yuan2 import Yuan2
-
- return Yuan2
-
-
-def _import_volcengine_maas() -> Type[BaseLLM]:
- from langchain_community.llms.volcengine_maas import VolcEngineMaasLLM
-
- return VolcEngineMaasLLM
-
-
-def _import_sparkllm() -> Type[BaseLLM]:
- from langchain_community.llms.sparkllm import SparkLLM
-
- return SparkLLM
-
-
-def _import_you() -> Type[BaseLLM]:
- from langchain_community.llms.you import You
-
- return You
-
-
-def _import_yi() -> Type[BaseLLM]:
- from langchain_community.llms.yi import YiLLM
-
- return YiLLM
-
-
-def __getattr__(name: str) -> Any:
- if name == "AI21":
- return _import_ai21()
- elif name == "AlephAlpha":
- return _import_aleph_alpha()
- elif name == "AmazonAPIGateway":
- return _import_amazon_api_gateway()
- elif name == "Anthropic":
- return _import_anthropic()
- elif name == "Anyscale":
- return _import_anyscale()
- elif name == "Aphrodite":
- return _import_aphrodite()
- elif name == "Arcee":
- return _import_arcee()
- elif name == "Aviary":
- return _import_aviary()
- elif name == "AzureMLOnlineEndpoint":
- return _import_azureml_endpoint()
- elif name == "BaichuanLLM" or name == "Baichuan":
- return _import_baichuan()
- elif name == "QianfanLLMEndpoint":
- return _import_baidu_qianfan_endpoint()
- elif name == "Banana":
- return _import_bananadev()
- elif name == "Baseten":
- return _import_baseten()
- elif name == "Beam":
- return _import_beam()
- elif name == "Bedrock":
- return _import_bedrock()
- elif name == "BigdlLLM":
- return _import_bigdlllm()
- elif name == "NIBittensorLLM":
- return _import_bittensor()
- elif name == "CerebriumAI":
- return _import_cerebriumai()
- elif name == "ChatGLM":
- return _import_chatglm()
- elif name == "Clarifai":
- return _import_clarifai()
- elif name == "Cohere":
- return _import_cohere()
- elif name == "CTransformers":
- return _import_ctransformers()
- elif name == "CTranslate2":
- return _import_ctranslate2()
- elif name == "Databricks":
- return _import_databricks()
- elif name == "DeepInfra":
- return _import_deepinfra()
- elif name == "DeepSparse":
- return _import_deepsparse()
- elif name == "EdenAI":
- return _import_edenai()
- elif name == "FakeListLLM":
- return _import_fake()
- elif name == "Fireworks":
- return _import_fireworks()
- elif name == "ForefrontAI":
- return _import_forefrontai()
- elif name == "Friendli":
- return _import_friendli()
- elif name == "GigaChat":
- return _import_gigachat()
- elif name == "GooglePalm":
- return _import_google_palm()
- elif name == "GooseAI":
- return _import_gooseai()
- elif name == "GPT4All":
- return _import_gpt4all()
- elif name == "GradientLLM":
- return _import_gradient_ai()
- elif name == "HuggingFaceEndpoint":
- return _import_huggingface_endpoint()
- elif name == "HuggingFaceHub":
- return _import_huggingface_hub()
- elif name == "HuggingFacePipeline":
- return _import_huggingface_pipeline()
- elif name == "HuggingFaceTextGenInference":
- return _import_huggingface_text_gen_inference()
- elif name == "HumanInputLLM":
- return _import_human()
- elif name == "IpexLLM":
- return _import_ipex_llm()
- elif name == "JavelinAIGateway":
- return _import_javelin_ai_gateway()
- elif name == "KoboldApiLLM":
- return _import_koboldai()
- elif name == "Konko":
- return _import_konko()
- elif name == "LlamaCpp":
- return _import_llamacpp()
- elif name == "Llamafile":
- return _import_llamafile()
- elif name == "ManifestWrapper":
- return _import_manifest()
- elif name == "Minimax":
- return _import_minimax()
- elif name == "Mlflow":
- return _import_mlflow()
- elif name == "MlflowAIGateway":
- return _import_mlflow_ai_gateway()
- elif name == "MLXPipeline":
- return _import_mlx_pipeline()
- elif name == "Modal":
- return _import_modal()
- elif name == "MosaicML":
- return _import_mosaicml()
- elif name == "NLPCloud":
- return _import_nlpcloud()
- elif name == "OCIModelDeploymentTGI":
- return _import_oci_md_tgi()
- elif name == "OCIModelDeploymentVLLM":
- return _import_oci_md_vllm()
- elif name == "OCIModelDeploymentLLM":
- return _import_oci_md()
- elif name == "OCIGenAI":
- return _import_oci_gen_ai()
- elif name == "OctoAIEndpoint":
- return _import_octoai_endpoint()
- elif name == "Ollama":
- return _import_ollama()
- elif name == "OpaquePrompts":
- return _import_opaqueprompts()
- elif name == "AzureOpenAI":
- return _import_azure_openai()
- elif name == "OpenAI":
- return _import_openai()
- elif name == "OpenAIChat":
- return _import_openai_chat()
- elif name == "OpenLLM":
- return _import_openllm()
- elif name == "OpenLM":
- return _import_openlm()
- elif name == "Outlines":
- return _import_outlines()
- elif name == "PaiEasEndpoint":
- return _import_pai_eas_endpoint()
- elif name == "Petals":
- return _import_petals()
- elif name == "PipelineAI":
- return _import_pipelineai()
- elif name == "Predibase":
- return _import_predibase()
- elif name == "PredictionGuard":
- return _import_predictionguard()
- elif name == "PromptLayerOpenAI":
- return _import_promptlayer()
- elif name == "PromptLayerOpenAIChat":
- return _import_promptlayer_chat()
- elif name == "Replicate":
- return _import_replicate()
- elif name == "RWKV":
- return _import_rwkv()
- elif name == "SagemakerEndpoint":
- return _import_sagemaker_endpoint()
- elif name == "SambaNovaCloud":
- return _import_sambanovacloud()
- elif name == "SambaStudio":
- return _import_sambastudio()
- elif name == "SelfHostedPipeline":
- return _import_self_hosted()
- elif name == "SelfHostedHuggingFaceLLM":
- return _import_self_hosted_hugging_face()
- elif name == "StochasticAI":
- return _import_stochasticai()
- elif name == "Nebula":
- return _import_symblai_nebula()
- elif name == "TextGen":
- return _import_textgen()
- elif name == "TitanTakeoff":
- return _import_titan_takeoff()
- elif name == "TitanTakeoffPro":
- return _import_titan_takeoff_pro()
- elif name == "Together":
- return _import_together()
- elif name == "Tongyi":
- return _import_tongyi()
- elif name == "VertexAI":
- return _import_vertex()
- elif name == "VertexAIModelGarden":
- return _import_vertex_model_garden()
- elif name == "VLLM":
- return _import_vllm()
- elif name == "VLLMOpenAI":
- return _import_vllm_openai()
- elif name == "WatsonxLLM":
- return _import_watsonxllm()
- elif name == "WeightOnlyQuantPipeline":
- return _import_weight_only_quantization()
- elif name == "Writer":
- return _import_writer()
- elif name == "Xinference":
- return _import_xinference()
- elif name == "YandexGPT":
- return _import_yandex_gpt()
- elif name == "Yuan2":
- return _import_yuan2()
- elif name == "VolcEngineMaasLLM":
- return _import_volcengine_maas()
- elif name == "SparkLLM":
- return _import_sparkllm()
- elif name == "YiLLM":
- return _import_yi()
- elif name == "You":
- return _import_you()
- elif name == "type_to_cls_dict":
- # for backwards compatibility
- type_to_cls_dict: Dict[str, Type[BaseLLM]] = {
- k: v() for k, v in get_type_to_cls_dict().items()
- }
- return type_to_cls_dict
- else:
- raise AttributeError(f"Could not find: {name}")
-
-
-__all__ = [
- "AI21",
- "AlephAlpha",
- "AmazonAPIGateway",
- "Anthropic",
- "Anyscale",
- "Aphrodite",
- "Arcee",
- "Aviary",
- "AzureMLOnlineEndpoint",
- "AzureOpenAI",
- "BaichuanLLM",
- "Banana",
- "Baseten",
- "Beam",
- "Bedrock",
- "CTransformers",
- "CTranslate2",
- "CerebriumAI",
- "ChatGLM",
- "Clarifai",
- "Cohere",
- "Databricks",
- "DeepInfra",
- "DeepSparse",
- "EdenAI",
- "FakeListLLM",
- "Fireworks",
- "ForefrontAI",
- "Friendli",
- "GPT4All",
- "GigaChat",
- "GooglePalm",
- "GooseAI",
- "GradientLLM",
- "HuggingFaceEndpoint",
- "HuggingFaceHub",
- "HuggingFacePipeline",
- "HuggingFaceTextGenInference",
- "HumanInputLLM",
- "IpexLLM",
- "JavelinAIGateway",
- "KoboldApiLLM",
- "Konko",
- "LlamaCpp",
- "Llamafile",
- "ManifestWrapper",
- "Minimax",
- "Mlflow",
- "MlflowAIGateway",
- "MLXPipeline",
- "Modal",
- "MosaicML",
- "NIBittensorLLM",
- "NLPCloud",
- "Nebula",
- "OCIGenAI",
- "OCIModelDeploymentTGI",
- "OCIModelDeploymentVLLM",
- "OCIModelDeploymentLLM",
- "OctoAIEndpoint",
- "Ollama",
- "OpaquePrompts",
- "OpenAI",
- "OpenAIChat",
- "OpenLLM",
- "OpenLM",
- "Outlines",
- "PaiEasEndpoint",
- "Petals",
- "PipelineAI",
- "Predibase",
- "PredictionGuard",
- "PromptLayerOpenAI",
- "PromptLayerOpenAIChat",
- "QianfanLLMEndpoint",
- "RWKV",
- "Replicate",
- "SagemakerEndpoint",
- "SambaNovaCloud",
- "SambaStudio",
- "SelfHostedHuggingFaceLLM",
- "SelfHostedPipeline",
- "SparkLLM",
- "StochasticAI",
- "TextGen",
- "TitanTakeoff",
- "TitanTakeoffPro",
- "Together",
- "Tongyi",
- "VLLM",
- "VLLMOpenAI",
- "VertexAI",
- "VertexAIModelGarden",
- "VolcEngineMaasLLM",
- "WatsonxLLM",
- "WeightOnlyQuantPipeline",
- "Writer",
- "Xinference",
- "YandexGPT",
- "Yuan2",
- "YiLLM",
- "You",
-]
-
-
-def get_type_to_cls_dict() -> Dict[str, Callable[[], Type[BaseLLM]]]:
- return {
- "ai21": _import_ai21,
- "aleph_alpha": _import_aleph_alpha,
- "amazon_api_gateway": _import_amazon_api_gateway,
- "amazon_bedrock": _import_bedrock,
- "anthropic": _import_anthropic,
- "anyscale": _import_anyscale,
- "arcee": _import_arcee,
- "aviary": _import_aviary,
- "azure": _import_azure_openai,
- "azureml_endpoint": _import_azureml_endpoint,
- "baichuan": _import_baichuan,
- "bananadev": _import_bananadev,
- "baseten": _import_baseten,
- "beam": _import_beam,
- "cerebriumai": _import_cerebriumai,
- "chat_glm": _import_chatglm,
- "clarifai": _import_clarifai,
- "cohere": _import_cohere,
- "ctransformers": _import_ctransformers,
- "ctranslate2": _import_ctranslate2,
- "databricks": _import_databricks,
- "databricks-chat": _import_databricks_chat, # deprecated / only for back compat
- "deepinfra": _import_deepinfra,
- "deepsparse": _import_deepsparse,
- "edenai": _import_edenai,
- "fake-list": _import_fake,
- "forefrontai": _import_forefrontai,
- "friendli": _import_friendli,
- "giga-chat-model": _import_gigachat,
- "google_palm": _import_google_palm,
- "gooseai": _import_gooseai,
- "gradient": _import_gradient_ai,
- "gpt4all": _import_gpt4all,
- "huggingface_endpoint": _import_huggingface_endpoint,
- "huggingface_hub": _import_huggingface_hub,
- "huggingface_pipeline": _import_huggingface_pipeline,
- "huggingface_textgen_inference": _import_huggingface_text_gen_inference,
- "human-input": _import_human,
- "koboldai": _import_koboldai,
- "konko": _import_konko,
- "llamacpp": _import_llamacpp,
- "llamafile": _import_llamafile,
- "textgen": _import_textgen,
- "minimax": _import_minimax,
- "mlflow": _import_mlflow,
- "mlflow-chat": _import_mlflow_chat, # deprecated / only for back compat
- "mlflow-ai-gateway": _import_mlflow_ai_gateway,
- "mlx_pipeline": _import_mlx_pipeline,
- "modal": _import_modal,
- "mosaic": _import_mosaicml,
- "nebula": _import_symblai_nebula,
- "nibittensor": _import_bittensor,
- "nlpcloud": _import_nlpcloud,
- "oci_model_deployment_tgi_endpoint": _import_oci_md_tgi,
- "oci_model_deployment_vllm_endpoint": _import_oci_md_vllm,
- "oci_model_deployment_endpoint": _import_oci_md,
- "oci_generative_ai": _import_oci_gen_ai,
- "octoai_endpoint": _import_octoai_endpoint,
- "ollama": _import_ollama,
- "openai": _import_openai,
- "openlm": _import_openlm,
- "pai_eas_endpoint": _import_pai_eas_endpoint,
- "petals": _import_petals,
- "pipelineai": _import_pipelineai,
- "predibase": _import_predibase,
- "opaqueprompts": _import_opaqueprompts,
- "replicate": _import_replicate,
- "rwkv": _import_rwkv,
- "sagemaker_endpoint": _import_sagemaker_endpoint,
- "sambanovacloud": _import_sambanovacloud,
- "sambastudio": _import_sambastudio,
- "self_hosted": _import_self_hosted,
- "self_hosted_hugging_face": _import_self_hosted_hugging_face,
- "stochasticai": _import_stochasticai,
- "together": _import_together,
- "tongyi": _import_tongyi,
- "titan_takeoff": _import_titan_takeoff,
- "titan_takeoff_pro": _import_titan_takeoff_pro,
- "vertexai": _import_vertex,
- "vertexai_model_garden": _import_vertex_model_garden,
- "openllm": _import_openllm,
- "outlines": _import_outlines,
- "vllm": _import_vllm,
- "vllm_openai": _import_vllm_openai,
- "watsonxllm": _import_watsonxllm,
- "weight_only_quantization": _import_weight_only_quantization,
- "writer": _import_writer,
- "xinference": _import_xinference,
- "javelin-ai-gateway": _import_javelin_ai_gateway,
- "qianfan_endpoint": _import_baidu_qianfan_endpoint,
- "yandex_gpt": _import_yandex_gpt,
- "yuan2": _import_yuan2,
- "VolcEngineMaasLLM": _import_volcengine_maas,
- "SparkLLM": _import_sparkllm,
- "yi": _import_yi,
- "you": _import_you,
- }
diff --git a/libs/community/langchain_community/llms/ai21.py b/libs/community/langchain_community/llms/ai21.py
deleted file mode 100644
index 08afd82a94..0000000000
--- a/libs/community/langchain_community/llms/ai21.py
+++ /dev/null
@@ -1,157 +0,0 @@
-from typing import Any, Dict, List, Optional, cast
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict, SecretStr
-
-
-class AI21PenaltyData(BaseModel):
- """Parameters for AI21 penalty data."""
-
- scale: int = 0
- applyToWhitespaces: bool = True
- applyToPunctuations: bool = True
- applyToNumbers: bool = True
- applyToStopwords: bool = True
- applyToEmojis: bool = True
-
-
-class AI21(LLM):
- """AI21 large language models.
-
- To use, you should have the environment variable ``AI21_API_KEY``
- set with your API key or pass it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import AI21
- ai21 = AI21(ai21_api_key="my-api-key", model="j2-jumbo-instruct")
- """
-
- model: str = "j2-jumbo-instruct"
- """Model name to use."""
-
- temperature: float = 0.7
- """What sampling temperature to use."""
-
- maxTokens: int = 256
- """The maximum number of tokens to generate in the completion."""
-
- minTokens: int = 0
- """The minimum number of tokens to generate in the completion."""
-
- topP: float = 1.0
- """Total probability mass of tokens to consider at each step."""
-
- presencePenalty: AI21PenaltyData = AI21PenaltyData()
- """Penalizes repeated tokens."""
-
- countPenalty: AI21PenaltyData = AI21PenaltyData()
- """Penalizes repeated tokens according to count."""
-
- frequencyPenalty: AI21PenaltyData = AI21PenaltyData()
- """Penalizes repeated tokens according to frequency."""
-
- numResults: int = 1
- """How many completions to generate for each prompt."""
-
- logitBias: Optional[Dict[str, float]] = None
- """Adjust the probability of specific tokens being generated."""
-
- ai21_api_key: Optional[SecretStr] = None
-
- stop: Optional[List[str]] = None
-
- base_url: Optional[str] = None
- """Base url to use, if None decides based on model name."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key exists in environment."""
- ai21_api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "ai21_api_key", "AI21_API_KEY")
- )
- values["ai21_api_key"] = ai21_api_key
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling AI21 API."""
- return {
- "temperature": self.temperature,
- "maxTokens": self.maxTokens,
- "minTokens": self.minTokens,
- "topP": self.topP,
- "presencePenalty": self.presencePenalty.dict(),
- "countPenalty": self.countPenalty.dict(),
- "frequencyPenalty": self.frequencyPenalty.dict(),
- "numResults": self.numResults,
- "logitBias": self.logitBias,
- }
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {**{"model": self.model}, **self._default_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "ai21"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to AI21's complete endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = ai21("Tell me a joke.")
- """
- if self.stop is not None and stop is not None:
- raise ValueError("`stop` found in both the input and default params.")
- elif self.stop is not None:
- stop = self.stop
- elif stop is None:
- stop = []
- if self.base_url is not None:
- base_url = self.base_url
- else:
- if self.model in ("j1-grande-instruct",):
- base_url = "https://api.ai21.com/studio/v1/experimental"
- else:
- base_url = "https://api.ai21.com/studio/v1"
- params = {**self._default_params, **kwargs}
- self.ai21_api_key = cast(SecretStr, self.ai21_api_key)
- response = requests.post(
- url=f"{base_url}/{self.model}/complete",
- headers={"Authorization": f"Bearer {self.ai21_api_key.get_secret_value()}"},
- json={"prompt": prompt, "stopSequences": stop, **params},
- )
- if response.status_code != 200:
- optional_detail = response.json().get("error")
- raise ValueError(
- f"AI21 /complete call failed with status code {response.status_code}."
- f" Details: {optional_detail}"
- )
- response_json = response.json()
- return response_json["completions"][0]["data"]["text"]
diff --git a/libs/community/langchain_community/llms/aleph_alpha.py b/libs/community/langchain_community/llms/aleph_alpha.py
deleted file mode 100644
index d1cf2175eb..0000000000
--- a/libs/community/langchain_community/llms/aleph_alpha.py
+++ /dev/null
@@ -1,286 +0,0 @@
-from typing import Any, Dict, List, Optional, Sequence
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import ConfigDict, SecretStr
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-
-class AlephAlpha(LLM):
- """Aleph Alpha large language models.
-
- To use, you should have the ``aleph_alpha_client`` python package installed, and the
- environment variable ``ALEPH_ALPHA_API_KEY`` set with your API key, or pass
- it as a named parameter to the constructor.
-
- Parameters are explained more in depth here:
- https://github.com/Aleph-Alpha/aleph-alpha-client/blob/c14b7dd2b4325c7da0d6a119f6e76385800e097b/aleph_alpha_client/completion.py#L10
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import AlephAlpha
- aleph_alpha = AlephAlpha(aleph_alpha_api_key="my-api-key")
- """
-
- client: Any = None #: :meta private:
- model: Optional[str] = "luminous-base"
- """Model name to use."""
-
- maximum_tokens: int = 64
- """The maximum number of tokens to be generated."""
-
- temperature: float = 0.0
- """A non-negative float that tunes the degree of randomness in generation."""
-
- top_k: int = 0
- """Number of most likely tokens to consider at each step."""
-
- top_p: float = 0.0
- """Total probability mass of tokens to consider at each step."""
-
- presence_penalty: float = 0.0
- """Penalizes repeated tokens."""
-
- frequency_penalty: float = 0.0
- """Penalizes repeated tokens according to frequency."""
-
- repetition_penalties_include_prompt: Optional[bool] = False
- """Flag deciding whether presence penalty or frequency penalty are
- updated from the prompt."""
-
- use_multiplicative_presence_penalty: Optional[bool] = False
- """Flag deciding whether presence penalty is applied
- multiplicatively (True) or additively (False)."""
-
- penalty_bias: Optional[str] = None
- """Penalty bias for the completion."""
-
- penalty_exceptions: Optional[List[str]] = None
- """List of strings that may be generated without penalty,
- regardless of other penalty settings"""
-
- penalty_exceptions_include_stop_sequences: Optional[bool] = None
- """Should stop_sequences be included in penalty_exceptions."""
-
- best_of: Optional[int] = None
- """returns the one with the "best of" results
- (highest log probability per token)
- """
-
- n: int = 1
- """How many completions to generate for each prompt."""
-
- logit_bias: Optional[Dict[int, float]] = None
- """The logit bias allows to influence the likelihood of generating tokens."""
-
- log_probs: Optional[int] = None
- """Number of top log probabilities to be returned for each generated token."""
-
- tokens: Optional[bool] = False
- """return tokens of completion."""
-
- disable_optimizations: Optional[bool] = False
-
- minimum_tokens: Optional[int] = 0
- """Generate at least this number of tokens."""
-
- echo: bool = False
- """Echo the prompt in the completion."""
-
- use_multiplicative_frequency_penalty: bool = False
-
- sequence_penalty: float = 0.0
-
- sequence_penalty_min_length: int = 2
-
- use_multiplicative_sequence_penalty: bool = False
-
- completion_bias_inclusion: Optional[Sequence[str]] = None
-
- completion_bias_inclusion_first_token_only: bool = False
-
- completion_bias_exclusion: Optional[Sequence[str]] = None
-
- completion_bias_exclusion_first_token_only: bool = False
- """Only consider the first token for the completion_bias_exclusion."""
-
- contextual_control_threshold: Optional[float] = None
- """If set to None, attention control parameters only apply to those tokens that have
- explicitly been set in the request.
- If set to a non-None value, control parameters are also applied to similar tokens.
- """
-
- control_log_additive: Optional[bool] = True
- """True: apply control by adding the log(control_factor) to attention scores.
- False: (attention_scores - - attention_scores.min(-1)) * control_factor
- """
-
- repetition_penalties_include_completion: bool = True
- """Flag deciding whether presence penalty or frequency penalty
- are updated from the completion."""
-
- raw_completion: bool = False
- """Force the raw completion of the model to be returned."""
-
- stop_sequences: Optional[List[str]] = None
- """Stop sequences to use."""
-
- # Client params
- aleph_alpha_api_key: Optional[SecretStr] = None
- """API key for Aleph Alpha API."""
- host: str = "https://api.aleph-alpha.com"
- """The hostname of the API host.
- The default one is "https://api.aleph-alpha.com")"""
- hosting: Optional[str] = None
- """Determines in which datacenters the request may be processed.
- You can either set the parameter to "aleph-alpha" or omit it (defaulting to None).
- Not setting this value, or setting it to None, gives us maximal
- flexibility in processing your request in our
- own datacenters and on servers hosted with other providers.
- Choose this option for maximal availability.
- Setting it to "aleph-alpha" allows us to only process the
- request in our own datacenters.
- Choose this option for maximal data privacy."""
- request_timeout_seconds: int = 305
- """Client timeout that will be set for HTTP requests in the
- `requests` library's API calls.
- Server will close all requests after 300 seconds with an internal server error."""
- total_retries: int = 8
- """The number of retries made in case requests fail with certain retryable
- status codes. If the last
- retry fails a corresponding exception is raised. Note, that between retries
- an exponential backoff
- is applied, starting with 0.5 s after the first retry and doubling for
- each retry made. So with the
- default setting of 8 retries a total wait time of 63.5 s is added
- between the retries."""
- nice: bool = False
- """Setting this to True, will signal to the API that you intend to be
- nice to other users
- by de-prioritizing your request below concurrent ones."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- values["aleph_alpha_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(values, "aleph_alpha_api_key", "ALEPH_ALPHA_API_KEY")
- )
- try:
- from aleph_alpha_client import Client
-
- values["client"] = Client(
- token=values["aleph_alpha_api_key"].get_secret_value(),
- host=values["host"],
- hosting=values["hosting"],
- request_timeout_seconds=values["request_timeout_seconds"],
- total_retries=values["total_retries"],
- nice=values["nice"],
- )
- except ImportError:
- raise ImportError(
- "Could not import aleph_alpha_client python package. "
- "Please install it with `pip install aleph_alpha_client`."
- )
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling the Aleph Alpha API."""
- return {
- "maximum_tokens": self.maximum_tokens,
- "temperature": self.temperature,
- "top_k": self.top_k,
- "top_p": self.top_p,
- "presence_penalty": self.presence_penalty,
- "frequency_penalty": self.frequency_penalty,
- "n": self.n,
- "repetition_penalties_include_prompt": self.repetition_penalties_include_prompt, # noqa: E501
- "use_multiplicative_presence_penalty": self.use_multiplicative_presence_penalty, # noqa: E501
- "penalty_bias": self.penalty_bias,
- "penalty_exceptions": self.penalty_exceptions,
- "penalty_exceptions_include_stop_sequences": self.penalty_exceptions_include_stop_sequences, # noqa: E501
- "best_of": self.best_of,
- "logit_bias": self.logit_bias,
- "log_probs": self.log_probs,
- "tokens": self.tokens,
- "disable_optimizations": self.disable_optimizations,
- "minimum_tokens": self.minimum_tokens,
- "echo": self.echo,
- "use_multiplicative_frequency_penalty": self.use_multiplicative_frequency_penalty, # noqa: E501
- "sequence_penalty": self.sequence_penalty,
- "sequence_penalty_min_length": self.sequence_penalty_min_length,
- "use_multiplicative_sequence_penalty": self.use_multiplicative_sequence_penalty, # noqa: E501
- "completion_bias_inclusion": self.completion_bias_inclusion,
- "completion_bias_inclusion_first_token_only": self.completion_bias_inclusion_first_token_only, # noqa: E501
- "completion_bias_exclusion": self.completion_bias_exclusion,
- "completion_bias_exclusion_first_token_only": self.completion_bias_exclusion_first_token_only, # noqa: E501
- "contextual_control_threshold": self.contextual_control_threshold,
- "control_log_additive": self.control_log_additive,
- "repetition_penalties_include_completion": self.repetition_penalties_include_completion, # noqa: E501
- "raw_completion": self.raw_completion,
- }
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {**{"model": self.model}, **self._default_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "aleph_alpha"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Aleph Alpha's completion endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = aleph_alpha("Tell me a joke.")
- """
- from aleph_alpha_client import CompletionRequest, Prompt
-
- params = self._default_params
- if self.stop_sequences is not None and stop is not None:
- raise ValueError(
- "stop sequences found in both the input and default params."
- )
- elif self.stop_sequences is not None:
- params["stop_sequences"] = self.stop_sequences
- else:
- params["stop_sequences"] = stop
- params = {**params, **kwargs}
- request = CompletionRequest(prompt=Prompt.from_text(prompt), **params)
- response = self.client.complete(model=self.model, request=request)
- text = response.completions[0].completion
- # If stop tokens are provided, Aleph Alpha's endpoint returns them.
- # In order to make this consistent with other endpoints, we strip them.
- if stop is not None or self.stop_sequences is not None:
- text = enforce_stop_tokens(text, params["stop_sequences"])
- return text
-
-
-if __name__ == "__main__":
- aa = AlephAlpha()
-
- print(aa.invoke("How are you?")) # noqa: T201
diff --git a/libs/community/langchain_community/llms/amazon_api_gateway.py b/libs/community/langchain_community/llms/amazon_api_gateway.py
deleted file mode 100644
index 61c088c8cd..0000000000
--- a/libs/community/langchain_community/llms/amazon_api_gateway.py
+++ /dev/null
@@ -1,103 +0,0 @@
-from typing import Any, Dict, List, Mapping, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from pydantic import ConfigDict
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-
-class ContentHandlerAmazonAPIGateway:
- """Adapter to prepare the inputs from Langchain to a format
- that LLM model expects.
-
- It also provides helper function to extract
- the generated text from the model response."""
-
- @classmethod
- def transform_input(
- cls, prompt: str, model_kwargs: Dict[str, Any]
- ) -> Dict[str, Any]:
- return {"inputs": prompt, "parameters": model_kwargs}
-
- @classmethod
- def transform_output(cls, response: Any) -> str:
- return response.json()[0]["generated_text"]
-
-
-class AmazonAPIGateway(LLM):
- """Amazon API Gateway to access LLM models hosted on AWS."""
-
- api_url: str
- """API Gateway URL"""
-
- headers: Optional[Dict] = None
- """API Gateway HTTP Headers to send, e.g. for authentication"""
-
- model_kwargs: Optional[Dict] = None
- """Keyword arguments to pass to the model."""
-
- content_handler: ContentHandlerAmazonAPIGateway = ContentHandlerAmazonAPIGateway()
- """The content handler class that provides an input and
- output transform functions to handle formats between LLM
- and the endpoint.
- """
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- **{"api_url": self.api_url, "headers": self.headers},
- **{"model_kwargs": _model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "amazon_api_gateway"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Amazon API Gateway model.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = se("Tell me a joke.")
- """
- _model_kwargs = self.model_kwargs or {}
- payload = self.content_handler.transform_input(prompt, _model_kwargs)
-
- try:
- response = requests.post(
- self.api_url,
- headers=self.headers,
- json=payload,
- )
- text = self.content_handler.transform_output(response)
-
- except Exception as error:
- raise ValueError(f"Error raised by the service: {error}")
-
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
-
- return text
diff --git a/libs/community/langchain_community/llms/anthropic.py b/libs/community/langchain_community/llms/anthropic.py
deleted file mode 100644
index 0a6af6799d..0000000000
--- a/libs/community/langchain_community/llms/anthropic.py
+++ /dev/null
@@ -1,360 +0,0 @@
-import re
-import warnings
-from typing import (
- Any,
- AsyncIterator,
- Callable,
- Dict,
- Iterator,
- List,
- Mapping,
- Optional,
-)
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models import BaseLanguageModel
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.prompt_values import PromptValue
-from langchain_core.utils import (
- check_package_version,
- get_from_dict_or_env,
- get_pydantic_field_names,
- pre_init,
-)
-from langchain_core.utils.utils import _build_model_kwargs, convert_to_secret_str
-from pydantic import ConfigDict, Field, SecretStr, model_validator
-
-
-class _AnthropicCommon(BaseLanguageModel):
- client: Any = None #: :meta private:
- async_client: Any = None #: :meta private:
- model: str = Field(default="claude-2", alias="model_name")
- """Model name to use."""
-
- max_tokens_to_sample: int = Field(default=256, alias="max_tokens")
- """Denotes the number of tokens to predict per generation."""
-
- temperature: Optional[float] = None
- """A non-negative float that tunes the degree of randomness in generation."""
-
- top_k: Optional[int] = None
- """Number of most likely tokens to consider at each step."""
-
- top_p: Optional[float] = None
- """Total probability mass of tokens to consider at each step."""
-
- streaming: bool = False
- """Whether to stream the results."""
-
- default_request_timeout: Optional[float] = None
- """Timeout for requests to Anthropic Completion API. Default is 600 seconds."""
-
- max_retries: int = 2
- """Number of retries allowed for requests sent to the Anthropic Completion API."""
-
- anthropic_api_url: Optional[str] = None
-
- anthropic_api_key: Optional[SecretStr] = None
-
- HUMAN_PROMPT: Optional[str] = None
- AI_PROMPT: Optional[str] = None
- count_tokens: Optional[Callable[[str], int]] = None
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict) -> Any:
- all_required_field_names = get_pydantic_field_names(cls)
- values = _build_model_kwargs(values, all_required_field_names)
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- values["anthropic_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(values, "anthropic_api_key", "ANTHROPIC_API_KEY")
- )
- # Get custom api url from environment.
- values["anthropic_api_url"] = get_from_dict_or_env(
- values,
- "anthropic_api_url",
- "ANTHROPIC_API_URL",
- default="https://api.anthropic.com",
- )
-
- try:
- import anthropic
-
- check_package_version("anthropic", gte_version="0.3")
- values["client"] = anthropic.Anthropic(
- base_url=values["anthropic_api_url"],
- api_key=values["anthropic_api_key"].get_secret_value(),
- timeout=values["default_request_timeout"],
- max_retries=values["max_retries"],
- )
- values["async_client"] = anthropic.AsyncAnthropic(
- base_url=values["anthropic_api_url"],
- api_key=values["anthropic_api_key"].get_secret_value(),
- timeout=values["default_request_timeout"],
- max_retries=values["max_retries"],
- )
- values["HUMAN_PROMPT"] = anthropic.HUMAN_PROMPT
- values["AI_PROMPT"] = anthropic.AI_PROMPT
- values["count_tokens"] = values["client"].count_tokens
-
- except ImportError:
- raise ImportError(
- "Could not import anthropic python package. "
- "Please it install it with `pip install anthropic`."
- )
- return values
-
- @property
- def _default_params(self) -> Mapping[str, Any]:
- """Get the default parameters for calling Anthropic API."""
- d = {
- "max_tokens_to_sample": self.max_tokens_to_sample,
- "model": self.model,
- }
- if self.temperature is not None:
- d["temperature"] = self.temperature
- if self.top_k is not None:
- d["top_k"] = self.top_k
- if self.top_p is not None:
- d["top_p"] = self.top_p
- return {**d, **self.model_kwargs}
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {**{}, **self._default_params}
-
- def _get_anthropic_stop(self, stop: Optional[List[str]] = None) -> List[str]:
- if not self.HUMAN_PROMPT or not self.AI_PROMPT:
- raise NameError("Please ensure the anthropic package is loaded")
-
- if stop is None:
- stop = []
-
- # Never want model to invent new turns of Human / Assistant dialog.
- stop.extend([self.HUMAN_PROMPT])
-
- return stop
-
-
-@deprecated(
- since="0.0.28",
- removal="1.0",
- alternative_import="langchain_anthropic.AnthropicLLM",
-)
-class Anthropic(LLM, _AnthropicCommon):
- """Anthropic large language models.
-
- To use, you should have the ``anthropic`` python package installed, and the
- environment variable ``ANTHROPIC_API_KEY`` set with your API key, or pass
- it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- import anthropic
- from langchain_community.llms import Anthropic
-
- model = Anthropic(model="", anthropic_api_key="my-api-key")
-
- # Simplest invocation, automatically wrapped with HUMAN_PROMPT
- # and AI_PROMPT.
- response = model.invoke("What are the biggest risks facing humanity?")
-
- # Or if you want to use the chat mode, build a few-shot-prompt, or
- # put words in the Assistant's mouth, use HUMAN_PROMPT and AI_PROMPT:
- raw_prompt = "What are the biggest risks facing humanity?"
- prompt = f"{anthropic.HUMAN_PROMPT} {prompt}{anthropic.AI_PROMPT}"
- response = model.invoke(prompt)
- """
-
- model_config = ConfigDict(
- populate_by_name=True,
- arbitrary_types_allowed=True,
- )
-
- @pre_init
- def raise_warning(cls, values: Dict) -> Dict:
- """Raise warning that this class is deprecated."""
- warnings.warn(
- "This Anthropic LLM is deprecated. "
- "Please use `from langchain_community.chat_models import ChatAnthropic` "
- "instead"
- )
- return values
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "anthropic-llm"
-
- def _wrap_prompt(self, prompt: str) -> str:
- if not self.HUMAN_PROMPT or not self.AI_PROMPT:
- raise NameError("Please ensure the anthropic package is loaded")
-
- if prompt.startswith(self.HUMAN_PROMPT):
- return prompt # Already wrapped.
-
- # Guard against common errors in specifying wrong number of newlines.
- corrected_prompt, n_subs = re.subn(r"^\n*Human:", self.HUMAN_PROMPT, prompt)
- if n_subs == 1:
- return corrected_prompt
-
- # As a last resort, wrap the prompt ourselves to emulate instruct-style.
- return f"{self.HUMAN_PROMPT} {prompt}{self.AI_PROMPT} Sure, here you go:\n"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- r"""Call out to Anthropic's completion endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- prompt = "What are the biggest risks facing humanity?"
- prompt = f"\n\nHuman: {prompt}\n\nAssistant:"
- response = model.invoke(prompt)
-
- """
- if self.streaming:
- completion = ""
- for chunk in self._stream(
- prompt=prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- completion += chunk.text
- return completion
-
- stop = self._get_anthropic_stop(stop)
- params = {**self._default_params, **kwargs}
- response = self.client.completions.create(
- prompt=self._wrap_prompt(prompt),
- stop_sequences=stop,
- **params,
- )
- return response.completion
-
- def convert_prompt(self, prompt: PromptValue) -> str:
- return self._wrap_prompt(prompt.to_string())
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Anthropic's completion endpoint asynchronously."""
- if self.streaming:
- completion = ""
- async for chunk in self._astream(
- prompt=prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- completion += chunk.text
- return completion
-
- stop = self._get_anthropic_stop(stop)
- params = {**self._default_params, **kwargs}
-
- response = await self.async_client.completions.create(
- prompt=self._wrap_prompt(prompt),
- stop_sequences=stop,
- **params,
- )
- return response.completion
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- r"""Call Anthropic completion_stream and return the resulting generator.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
- Returns:
- A generator representing the stream of tokens from Anthropic.
- Example:
- .. code-block:: python
-
- prompt = "Write a poem about a stream."
- prompt = f"\n\nHuman: {prompt}\n\nAssistant:"
- generator = anthropic.stream(prompt)
- for token in generator:
- yield token
- """
- stop = self._get_anthropic_stop(stop)
- params = {**self._default_params, **kwargs}
-
- for token in self.client.completions.create(
- prompt=self._wrap_prompt(prompt), stop_sequences=stop, stream=True, **params
- ):
- chunk = GenerationChunk(text=token.completion)
- if run_manager:
- run_manager.on_llm_new_token(chunk.text, chunk=chunk)
- yield chunk
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- r"""Call Anthropic completion_stream and return the resulting generator.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
- Returns:
- A generator representing the stream of tokens from Anthropic.
- Example:
- .. code-block:: python
- prompt = "Write a poem about a stream."
- prompt = f"\n\nHuman: {prompt}\n\nAssistant:"
- generator = anthropic.stream(prompt)
- for token in generator:
- yield token
- """
- stop = self._get_anthropic_stop(stop)
- params = {**self._default_params, **kwargs}
-
- async for token in await self.async_client.completions.create(
- prompt=self._wrap_prompt(prompt),
- stop_sequences=stop,
- stream=True,
- **params,
- ):
- chunk = GenerationChunk(text=token.completion)
- if run_manager:
- await run_manager.on_llm_new_token(chunk.text, chunk=chunk)
- yield chunk
-
- def get_num_tokens(self, text: str) -> int:
- """Calculate number of tokens."""
- if not self.count_tokens:
- raise NameError("Please ensure the anthropic package is loaded")
- return self.count_tokens(text)
diff --git a/libs/community/langchain_community/llms/anyscale.py b/libs/community/langchain_community/llms/anyscale.py
deleted file mode 100644
index b55fb68ca8..0000000000
--- a/libs/community/langchain_community/llms/anyscale.py
+++ /dev/null
@@ -1,319 +0,0 @@
-"""Wrapper around Anyscale Endpoint"""
-
-from typing import (
- Any,
- Dict,
- List,
- Mapping,
- Optional,
- Set,
-)
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.outputs import Generation, GenerationChunk, LLMResult
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import Field, SecretStr
-
-from langchain_community.llms.openai import (
- BaseOpenAI,
- acompletion_with_retry,
- completion_with_retry,
-)
-from langchain_community.utils.openai import is_openai_v1
-
-DEFAULT_BASE_URL = "https://api.endpoints.anyscale.com/v1"
-DEFAULT_MODEL = "mistralai/Mixtral-8x7B-Instruct-v0.1"
-
-
-def update_token_usage(
- keys: Set[str], response: Dict[str, Any], token_usage: Dict[str, Any]
-) -> None:
- """Update token usage."""
- _keys_to_use = keys.intersection(response["usage"])
- for _key in _keys_to_use:
- if _key not in token_usage:
- token_usage[_key] = response["usage"][_key]
- else:
- token_usage[_key] += response["usage"][_key]
-
-
-def create_llm_result(
- choices: Any, prompts: List[str], token_usage: Dict[str, int], model_name: str
-) -> LLMResult:
- """Create the LLMResult from the choices and prompts."""
- generations = []
- for i, _ in enumerate(prompts):
- choice = choices[i]
- generations.append(
- [
- Generation(
- text=choice["message"]["content"],
- generation_info=dict(
- finish_reason=choice.get("finish_reason"),
- logprobs=choice.get("logprobs"),
- ),
- )
- ]
- )
- llm_output = {"token_usage": token_usage, "model_name": model_name}
- return LLMResult(generations=generations, llm_output=llm_output)
-
-
-class Anyscale(BaseOpenAI):
- """Anyscale large language models.
-
- To use, you should have the environment variable ``ANYSCALE_API_KEY``set with your
- Anyscale Endpoint, or pass it as a named parameter to the constructor.
- To use with Anyscale Private Endpoint, please also set ``ANYSCALE_BASE_URL``.
-
- Example:
- .. code-block:: python
- from langchain.llms import Anyscale
- anyscalellm = Anyscale(anyscale_api_key="ANYSCALE_API_KEY")
- # To leverage Ray for parallel processing
- @ray.remote(num_cpus=1)
- def send_query(llm, text):
- resp = llm.invoke(text)
- return resp
- futures = [send_query.remote(anyscalellm, text) for text in texts]
- results = ray.get(futures)
- """
-
- """Key word arguments to pass to the model."""
- anyscale_api_base: str = Field(default=DEFAULT_BASE_URL)
- anyscale_api_key: SecretStr = Field(default=SecretStr(""))
- model_name: str = Field(default=DEFAULT_MODEL)
-
- prefix_messages: List = Field(default_factory=list)
-
- @classmethod
- def is_lc_serializable(cls) -> bool:
- return False
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- values["anyscale_api_base"] = get_from_dict_or_env(
- values,
- "anyscale_api_base",
- "ANYSCALE_API_BASE",
- default=DEFAULT_BASE_URL,
- )
- values["anyscale_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(values, "anyscale_api_key", "ANYSCALE_API_KEY")
- )
- values["model_name"] = get_from_dict_or_env(
- values,
- "model_name",
- "MODEL_NAME",
- default=DEFAULT_MODEL,
- )
-
- try:
- import openai
-
- if is_openai_v1():
- client_params = {
- "api_key": values["anyscale_api_key"].get_secret_value(),
- "base_url": values["anyscale_api_base"],
- # To do: future support
- # "organization": values["openai_organization"],
- # "timeout": values["request_timeout"],
- # "max_retries": values["max_retries"],
- # "default_headers": values["default_headers"],
- # "default_query": values["default_query"],
- # "http_client": values["http_client"],
- }
- if not values.get("client"):
- values["client"] = openai.OpenAI(**client_params).completions
- if not values.get("async_client"):
- values["async_client"] = openai.AsyncOpenAI(
- **client_params
- ).completions
- else:
- values["openai_api_base"] = values["anyscale_api_base"]
- values["openai_api_key"] = values["anyscale_api_key"].get_secret_value()
- values["client"] = openai.Completion
- except ImportError:
- raise ImportError(
- "Could not import openai python package. "
- "Please install it with `pip install openai`."
- )
- if values["streaming"] and values["n"] > 1:
- raise ValueError("Cannot stream results when n > 1.")
- if values["streaming"] and values["best_of"] > 1:
- raise ValueError("Cannot stream results when best_of > 1.")
-
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- **{"model_name": self.model_name},
- **super()._identifying_params,
- }
-
- @property
- def _invocation_params(self) -> Dict[str, Any]:
- """Get the parameters used to invoke the model."""
- openai_creds: Dict[str, Any] = {
- "model": self.model_name,
- }
- if not is_openai_v1():
- openai_creds.update(
- {
- "api_key": self.anyscale_api_key.get_secret_value(),
- "api_base": self.anyscale_api_base,
- }
- )
- return {**openai_creds, **super()._invocation_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "Anyscale LLM"
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call out to OpenAI's endpoint with k unique prompts.
-
- Args:
- prompts: The prompts to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The full LLM output.
-
- Example:
- .. code-block:: python
-
- response = openai.generate(["Tell me a joke."])
- """
- # TODO: write a unit test for this
- params = self._invocation_params
- params = {**params, **kwargs}
- sub_prompts = self.get_sub_prompts(params, prompts, stop)
- choices = []
- token_usage: Dict[str, int] = {}
- # Get the token usage from the response.
- # Includes prompt, completion, and total tokens used.
- _keys = {"completion_tokens", "prompt_tokens", "total_tokens"}
- system_fingerprint: Optional[str] = None
- for _prompts in sub_prompts:
- if self.streaming:
- if len(_prompts) > 1:
- raise ValueError("Cannot stream results with multiple prompts.")
-
- generation: Optional[GenerationChunk] = None
- for chunk in self._stream(_prompts[0], stop, run_manager, **kwargs):
- if generation is None:
- generation = chunk
- else:
- generation += chunk
- assert generation is not None
- choices.append(
- {
- "text": generation.text,
- "finish_reason": generation.generation_info.get("finish_reason")
- if generation.generation_info
- else None,
- "logprobs": generation.generation_info.get("logprobs")
- if generation.generation_info
- else None,
- }
- )
- else:
- response = completion_with_retry(
- ## THis is the ONLY change from BaseOpenAI()._generate()
- self,
- prompt=_prompts[0],
- run_manager=run_manager,
- **params,
- )
- if not isinstance(response, dict):
- # V1 client returns the response in an PyDantic object instead of
- # dict. For the transition period, we deep convert it to dict.
- response = response.dict()
-
- choices.extend(response["choices"])
- update_token_usage(_keys, response, token_usage)
- if not system_fingerprint:
- system_fingerprint = response.get("system_fingerprint")
- return self.create_llm_result(
- choices,
- prompts,
- params,
- token_usage,
- system_fingerprint=system_fingerprint,
- )
-
- async def _agenerate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call out to OpenAI's endpoint async with k unique prompts."""
- params = self._invocation_params
- params = {**params, **kwargs}
- sub_prompts = self.get_sub_prompts(params, prompts, stop)
- choices = []
- token_usage: Dict[str, int] = {}
- # Get the token usage from the response.
- # Includes prompt, completion, and total tokens used.
- _keys = {"completion_tokens", "prompt_tokens", "total_tokens"}
- system_fingerprint: Optional[str] = None
- for _prompts in sub_prompts:
- if self.streaming:
- if len(_prompts) > 1:
- raise ValueError("Cannot stream results with multiple prompts.")
-
- generation: Optional[GenerationChunk] = None
- async for chunk in self._astream(
- _prompts[0], stop, run_manager, **kwargs
- ):
- if generation is None:
- generation = chunk
- else:
- generation += chunk
- assert generation is not None
- choices.append(
- {
- "text": generation.text,
- "finish_reason": generation.generation_info.get("finish_reason")
- if generation.generation_info
- else None,
- "logprobs": generation.generation_info.get("logprobs")
- if generation.generation_info
- else None,
- }
- )
- else:
- response = await acompletion_with_retry(
- ## THis is the ONLY change from BaseOpenAI()._agenerate()
- self,
- prompt=_prompts[0],
- run_manager=run_manager,
- **params,
- )
- if not isinstance(response, dict):
- response = response.dict()
- choices.extend(response["choices"])
- update_token_usage(_keys, response, token_usage)
- return self.create_llm_result(
- choices,
- prompts,
- params,
- token_usage,
- system_fingerprint=system_fingerprint,
- )
diff --git a/libs/community/langchain_community/llms/aphrodite.py b/libs/community/langchain_community/llms/aphrodite.py
deleted file mode 100644
index bcaeabe296..0000000000
--- a/libs/community/langchain_community/llms/aphrodite.py
+++ /dev/null
@@ -1,251 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models import BaseLLM
-from langchain_core.outputs import Generation, LLMResult
-from langchain_core.utils import pre_init
-from pydantic import Field
-
-
-class Aphrodite(BaseLLM):
- """Aphrodite language model."""
-
- model: str = ""
- """The name or path of a HuggingFace Transformers model."""
-
- tensor_parallel_size: Optional[int] = 1
- """The number of GPUs to use for distributed execution with tensor parallelism."""
-
- trust_remote_code: Optional[bool] = False
- """Trust remote code (e.g., from HuggingFace) when downloading the model
- and tokenizer."""
-
- n: int = 1
- """Number of output sequences to return for the given prompt."""
-
- best_of: Optional[int] = None
- """Number of output sequences that are generated from the prompt.
- From these `best_of` sequences, the top `n` sequences are returned.
- `best_of` must be >= `n`. This is treated as the beam width when
- `use_beam_search` is True. By default, `best_of` is set to `n`."""
-
- presence_penalty: float = 0.0
- """Float that penalizes new tokens based on whether they appear in the
- generated text so far. Values > 0 encourage the model to generate new
- tokens, while values < 0 encourage the model to repeat tokens."""
-
- frequency_penalty: float = 0.0
- """Float that penalizes new tokens based on their frequency in the
- generated text so far. Applied additively to the logits."""
-
- repetition_penalty: float = 1.0
- """Float that penalizes new tokens based on their frequency in the
- generated text so far. Applied multiplicatively to the logits."""
-
- temperature: float = 1.0
- """Float that controls the randomness of the sampling. Lower values
- make the model more deterministic, while higher values make the model
- more random. Zero is equivalent to greedy sampling."""
-
- top_p: float = 1.0
- """Float that controls the cumulative probability of the top tokens to consider.
- Must be in (0, 1]. Set to 1.0 to consider all tokens."""
-
- top_k: int = -1
- """Integer that controls the number of top tokens to consider. Set to -1 to
- consider all tokens (disabled)."""
-
- top_a: float = 0.0
- """Float that controls the cutoff for Top-A sampling. Exact cutoff is
- top_a*max_prob**2. Must be in [0,inf], 0 to disable."""
-
- min_p: float = 0.0
- """Float that controls the cutoff for min-p sampling. Exact cutoff is
- min_p*max_prob. Must be in [0,1], 0 to disable."""
-
- tfs: float = 1.0
- """Float that controls the cumulative approximate curvature of the
- distribution to retain for Tail Free Sampling. Must be in (0, 1].
- Set to 1.0 to disable."""
-
- eta_cutoff: float = 0.0
- """Float that controls the cutoff threshold for Eta sampling
- (a form of entropy adaptive truncation sampling). Threshold is
- calculated as `min(eta, sqrt(eta)*entropy(probs)). Specified
- in units of 1e-4. Set to 0 to disable."""
-
- epsilon_cutoff: float = 0.0
- """Float that controls the cutoff threshold for Epsilon sampling
- (simple probability threshold truncation). Specified in units of
- 1e-4. Set to 0 to disable."""
-
- typical_p: float = 1.0
- """Float that controls the cumulative probability of tokens closest
- in surprise to the expected surprise to consider. Must be in (0, 1].
- Set to 1 to disable."""
-
- mirostat_mode: int = 0
- """The mirostat mode to use. 0 for no mirostat, 2 for mirostat v2.
- Mode 1 is not supported."""
-
- mirostat_tau: float = 0.0
- """The target 'surprisal' that mirostat works towards. Range [0, inf)."""
-
- use_beam_search: bool = False
- """Whether to use beam search instead of sampling."""
-
- length_penalty: float = 1.0
- """Float that penalizes sequences based on their length. Used only
- when `use_beam_search` is True."""
-
- early_stopping: bool = False
- """Controls the stopping condition for beam search. It accepts the
- following values: `True`, where the generation stops as soon as there
- are `best_of` complete candidates; `False`, where a heuristic is applied
- to the generation stops when it is very unlikely to find better candidates;
- `never`, where the beam search procedure only stops where there cannot be
- better candidates (canonical beam search algorithm)."""
-
- stop: Optional[List[str]] = None
- """List of strings that stop the generation when they are generated.
- The returned output will not contain the stop tokens."""
-
- stop_token_ids: Optional[List[int]] = None
- """List of tokens that stop the generation when they are generated.
- The returned output will contain the stop tokens unless the stop tokens
- are special tokens."""
-
- ignore_eos: bool = False
- """Whether to ignore the EOS token and continue generating tokens after
- the EOS token is generated."""
-
- max_tokens: int = 512
- """Maximum number of tokens to generate per output sequence."""
-
- logprobs: Optional[int] = None
- """Number of log probabilities to return per output token."""
-
- prompt_logprobs: Optional[int] = None
- """Number of log probabilities to return per prompt token."""
-
- custom_token_bans: Optional[List[int]] = None
- """List of token IDs to ban from generating."""
-
- skip_special_tokens: bool = True
- """Whether to skip special tokens in the output. Defaults to True."""
-
- spaces_between_special_tokens: bool = True
- """Whether to add spaces between special tokens in the output.
- Defaults to True."""
-
- logit_bias: Optional[Dict[str, float]] = None
- """List of LogitsProcessors to change the probability of token
- prediction at runtime."""
-
- dtype: str = "auto"
- """The data type for the model weights and activations."""
-
- download_dir: Optional[str] = None
- """Directory to download and load the weights. (Default to the default
- cache dir of huggingface)"""
-
- quantization: Optional[str] = None
- """Quantization mode to use. Can be one of `awq` or `gptq`."""
-
- aphrodite_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `aphrodite.LLM` call not explicitly
- specified."""
-
- client: Any = None #: :meta private:
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that python package exists in environment."""
-
- try:
- from aphrodite import LLM as AphroditeModel
- except ImportError:
- raise ImportError(
- "Could not import aphrodite-engine python package. "
- "Please install it with `pip install aphrodite-engine`."
- )
-
- # aphrodite_kwargs = values["aphrodite_kwargs"]
- # if values.get("quantization"):
- # aphrodite_kwargs["quantization"] = values["quantization"]
-
- values["client"] = AphroditeModel(
- model=values["model"],
- tensor_parallel_size=values["tensor_parallel_size"],
- trust_remote_code=values["trust_remote_code"],
- dtype=values["dtype"],
- download_dir=values["download_dir"],
- **values["aphrodite_kwargs"],
- )
-
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling aphrodite."""
- return {
- "n": self.n,
- "best_of": self.best_of,
- "max_tokens": self.max_tokens,
- "top_k": self.top_k,
- "top_p": self.top_p,
- "top_a": self.top_a,
- "min_p": self.min_p,
- "temperature": self.temperature,
- "presence_penalty": self.presence_penalty,
- "frequency_penalty": self.frequency_penalty,
- "repetition_penalty": self.repetition_penalty,
- "tfs": self.tfs,
- "eta_cutoff": self.eta_cutoff,
- "epsilon_cutoff": self.epsilon_cutoff,
- "typical_p": self.typical_p,
- "mirostat_mode": self.mirostat_mode,
- "mirostat_tau": self.mirostat_tau,
- "length_penalty": self.length_penalty,
- "early_stopping": self.early_stopping,
- "use_beam_search": self.use_beam_search,
- "stop": self.stop,
- "ignore_eos": self.ignore_eos,
- "logprobs": self.logprobs,
- "prompt_logprobs": self.prompt_logprobs,
- "custom_token_bans": self.custom_token_bans,
- "skip_special_tokens": self.skip_special_tokens,
- "spaces_between_special_tokens": self.spaces_between_special_tokens,
- "logit_bias": self.logit_bias,
- }
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Run the LLM on the given prompt and input."""
-
- from aphrodite import SamplingParams
-
- # build sampling parameters
- params = {**self._default_params, **kwargs, "stop": stop}
- if "logit_bias" in params:
- del params["logit_bias"]
- sampling_params = SamplingParams(**params)
- # call the model
- outputs = self.client.generate(prompts, sampling_params)
-
- generations = []
- for output in outputs:
- text = output.outputs[0].text
- generations.append([Generation(text=text)])
-
- return LLMResult(generations=generations)
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "aphrodite"
diff --git a/libs/community/langchain_community/llms/arcee.py b/libs/community/langchain_community/llms/arcee.py
deleted file mode 100644
index 42fcef87bb..0000000000
--- a/libs/community/langchain_community/llms/arcee.py
+++ /dev/null
@@ -1,146 +0,0 @@
-from typing import Any, Dict, List, Optional, Union, cast
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env
-from pydantic import ConfigDict, SecretStr, model_validator
-
-from langchain_community.utilities.arcee import ArceeWrapper, DALMFilter
-
-
-class Arcee(LLM):
- """Arcee's Domain Adapted Language Models (DALMs).
-
- To use, set the ``ARCEE_API_KEY`` environment variable with your Arcee API key,
- or pass ``arcee_api_key`` as a named parameter.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import Arcee
-
- arcee = Arcee(
- model="DALM-PubMed",
- arcee_api_key="ARCEE-API-KEY"
- )
-
- response = arcee("AI-driven music therapy")
- """
-
- _client: Optional[ArceeWrapper] = None #: :meta private:
- """Arcee _client."""
-
- arcee_api_key: Union[SecretStr, str, None] = None
- """Arcee API Key"""
-
- model: str
- """Arcee DALM name"""
-
- arcee_api_url: str = "https://api.arcee.ai"
- """Arcee API URL"""
-
- arcee_api_version: str = "v2"
- """Arcee API Version"""
-
- arcee_app_url: str = "https://app.arcee.ai"
- """Arcee App URL"""
-
- model_id: str = ""
- """Arcee Model ID"""
-
- model_kwargs: Optional[Dict[str, Any]] = None
- """Keyword arguments to pass to the model."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "arcee"
-
- def __init__(self, **data: Any) -> None:
- """Initializes private fields."""
-
- super().__init__(**data)
- api_key = cast(SecretStr, self.arcee_api_key)
- self._client = ArceeWrapper(
- arcee_api_key=api_key,
- arcee_api_url=self.arcee_api_url,
- arcee_api_version=self.arcee_api_version,
- model_kwargs=self.model_kwargs,
- model_name=self.model,
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environments(cls, values: Dict) -> Any:
- """Validate Arcee environment variables."""
-
- # validate env vars
- values["arcee_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(
- values,
- "arcee_api_key",
- "ARCEE_API_KEY",
- )
- )
-
- values["arcee_api_url"] = get_from_dict_or_env(
- values,
- "arcee_api_url",
- "ARCEE_API_URL",
- )
-
- values["arcee_app_url"] = get_from_dict_or_env(
- values,
- "arcee_app_url",
- "ARCEE_APP_URL",
- )
-
- values["arcee_api_version"] = get_from_dict_or_env(
- values,
- "arcee_api_version",
- "ARCEE_API_VERSION",
- )
-
- # validate model kwargs
- if values.get("model_kwargs"):
- kw = values["model_kwargs"]
-
- # validate size
- if kw.get("size") is not None:
- if not kw.get("size") >= 0:
- raise ValueError("`size` must be positive")
-
- # validate filters
- if kw.get("filters") is not None:
- if not isinstance(kw.get("filters"), List):
- raise ValueError("`filters` must be a list")
- for f in kw.get("filters"):
- DALMFilter(**f)
- return values
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Generate text from Arcee DALM.
-
- Args:
- prompt: Prompt to generate text from.
- size: The max number of context results to retrieve.
- Defaults to 3. (Can be less if filters are provided).
- filters: Filters to apply to the context dataset.
- """
-
- try:
- if not self._client:
- raise ValueError("Client is not initialized.")
- return self._client.generate(prompt=prompt, **kwargs)
- except Exception as e:
- raise Exception(f"Failed to generate text: {e}") from e
diff --git a/libs/community/langchain_community/llms/aviary.py b/libs/community/langchain_community/llms/aviary.py
deleted file mode 100644
index 95fd5730fa..0000000000
--- a/libs/community/langchain_community/llms/aviary.py
+++ /dev/null
@@ -1,197 +0,0 @@
-import dataclasses
-import os
-from typing import Any, Dict, List, Mapping, Optional, Union, cast
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import ConfigDict, model_validator
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-TIMEOUT = 60
-
-
-@dataclasses.dataclass
-class AviaryBackend:
- """Aviary backend.
-
- Attributes:
- backend_url: The URL for the Aviary backend.
- bearer: The bearer token for the Aviary backend.
- """
-
- backend_url: str
- bearer: str
-
- def __post_init__(self) -> None:
- self.header = {"Authorization": self.bearer}
-
- @classmethod
- def from_env(cls) -> "AviaryBackend":
- aviary_url = os.getenv("AVIARY_URL")
- assert aviary_url, "AVIARY_URL must be set"
-
- aviary_token = os.getenv("AVIARY_TOKEN", "")
-
- bearer = f"Bearer {aviary_token}" if aviary_token else ""
- aviary_url += "/" if not aviary_url.endswith("/") else ""
-
- return cls(aviary_url, bearer)
-
-
-def get_models() -> List[str]:
- """List available models"""
- backend = AviaryBackend.from_env()
- request_url = backend.backend_url + "-/routes"
- response = requests.get(request_url, headers=backend.header, timeout=TIMEOUT)
- try:
- result = response.json()
- except requests.JSONDecodeError as e:
- raise RuntimeError(
- f"Error decoding JSON from {request_url}. Text response: {response.text}"
- ) from e
- result = sorted(
- [k.lstrip("/").replace("--", "/") for k in result.keys() if "--" in k]
- )
- return result
-
-
-def get_completions(
- model: str,
- prompt: str,
- use_prompt_format: bool = True,
- version: str = "",
-) -> Dict[str, Union[str, float, int]]:
- """Get completions from Aviary models."""
-
- backend = AviaryBackend.from_env()
- url = backend.backend_url + model.replace("/", "--") + "/" + version + "query"
- response = requests.post(
- url,
- headers=backend.header,
- json={"prompt": prompt, "use_prompt_format": use_prompt_format},
- timeout=TIMEOUT,
- )
- try:
- return response.json()
- except requests.JSONDecodeError as e:
- raise RuntimeError(
- f"Error decoding JSON from {url}. Text response: {response.text}"
- ) from e
-
-
-class Aviary(LLM):
- """Aviary hosted models.
-
- Aviary is a backend for hosted models. You can
- find out more about aviary at
- http://github.com/ray-project/aviary
-
- To get a list of the models supported on an
- aviary, follow the instructions on the website to
- install the aviary CLI and then use:
- `aviary models`
-
- AVIARY_URL and AVIARY_TOKEN environment variables must be set.
-
- Attributes:
- model: The name of the model to use. Defaults to "amazon/LightGPT".
- aviary_url: The URL for the Aviary backend. Defaults to None.
- aviary_token: The bearer token for the Aviary backend. Defaults to None.
- use_prompt_format: If True, the prompt template for the model will be ignored.
- Defaults to True.
- version: API version to use for Aviary. Defaults to None.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import Aviary
- os.environ["AVIARY_URL"] = ""
- os.environ["AVIARY_TOKEN"] = ""
- light = Aviary(model='amazon/LightGPT')
- output = light('How do you make fried rice?')
- """
-
- model: str = "amazon/LightGPT"
- aviary_url: Optional[str] = None
- aviary_token: Optional[str] = None
- # If True the prompt template for the model will be ignored.
- use_prompt_format: bool = True
- # API version to use for Aviary
- version: Optional[str] = None
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and python package exists in environment."""
- aviary_url = get_from_dict_or_env(values, "aviary_url", "AVIARY_URL")
- aviary_token = get_from_dict_or_env(values, "aviary_token", "AVIARY_TOKEN")
-
- # Set env viarables for aviary sdk
- os.environ["AVIARY_URL"] = aviary_url
- os.environ["AVIARY_TOKEN"] = aviary_token
-
- try:
- aviary_models = get_models()
- except requests.exceptions.RequestException as e:
- raise ValueError(e)
-
- model = values.get("model")
- if model and model not in aviary_models:
- raise ValueError(f"{aviary_url} does not support model {values['model']}.")
-
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- "model_name": self.model,
- "aviary_url": self.aviary_url,
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return f"aviary-{self.model.replace('/', '-')}"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Aviary
- Args:
- prompt: The prompt to pass into the model.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = aviary("Tell me a joke.")
- """
- kwargs = {"use_prompt_format": self.use_prompt_format}
- if self.version:
- kwargs["version"] = self.version
-
- output = get_completions(
- model=self.model,
- prompt=prompt,
- **kwargs,
- )
-
- text = cast(str, output["generated_text"])
- if stop:
- text = enforce_stop_tokens(text, stop)
-
- return text
diff --git a/libs/community/langchain_community/llms/azureml_endpoint.py b/libs/community/langchain_community/llms/azureml_endpoint.py
deleted file mode 100644
index 184bf5f53e..0000000000
--- a/libs/community/langchain_community/llms/azureml_endpoint.py
+++ /dev/null
@@ -1,562 +0,0 @@
-import json
-import urllib.request
-import warnings
-from abc import abstractmethod
-from enum import Enum
-from typing import Any, Dict, List, Mapping, Optional
-
-from langchain_core.callbacks.manager import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import BaseLLM
-from langchain_core.outputs import Generation, LLMResult
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env
-from pydantic import BaseModel, ConfigDict, SecretStr, model_validator, validator
-
-DEFAULT_TIMEOUT = 50
-
-
-class AzureMLEndpointClient(object):
- """AzureML Managed Endpoint client."""
-
- def __init__(
- self,
- endpoint_url: str,
- endpoint_api_key: str,
- deployment_name: str = "",
- timeout: int = DEFAULT_TIMEOUT,
- ) -> None:
- """Initialize the class."""
- if not endpoint_api_key or not endpoint_url:
- raise ValueError(
- """A key/token and REST endpoint should
- be provided to invoke the endpoint"""
- )
- self.endpoint_url = endpoint_url
- self.endpoint_api_key = endpoint_api_key
- self.deployment_name = deployment_name
- self.timeout = timeout
-
- def call(
- self,
- body: bytes,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> bytes:
- """call."""
-
- # The azureml-model-deployment header will force the request to go to a
- # specific deployment. Remove this header to have the request observe the
- # endpoint traffic rules.
- headers = {
- "Content-Type": "application/json",
- "Authorization": ("Bearer " + self.endpoint_api_key),
- }
- if self.deployment_name != "":
- headers["azureml-model-deployment"] = self.deployment_name
-
- req = urllib.request.Request(self.endpoint_url, body, headers)
- response = urllib.request.urlopen(
- req, timeout=kwargs.get("timeout", self.timeout)
- )
- result = response.read()
- return result
-
-
-class AzureMLEndpointApiType(str, Enum):
- """Azure ML endpoints API types. Use `dedicated` for models deployed in hosted
- infrastructure (also known as Online Endpoints in Azure Machine Learning),
- or `serverless` for models deployed as a service with a
- pay-as-you-go billing or PTU.
- """
-
- dedicated = "dedicated"
- realtime = "realtime" #: Deprecated
- serverless = "serverless"
-
-
-class ContentFormatterBase:
- """Transform request and response of AzureML endpoint to match with
- required schema.
- """
-
- """
- Example:
- .. code-block:: python
-
- class ContentFormatter(ContentFormatterBase):
- content_type = "application/json"
- accepts = "application/json"
-
- def format_request_payload(
- self,
- prompt: str,
- model_kwargs: Dict,
- api_type: AzureMLEndpointApiType,
- ) -> bytes:
- input_str = json.dumps(
- {
- "inputs": {"input_string": [prompt]},
- "parameters": model_kwargs,
- }
- )
- return str.encode(input_str)
-
- def format_response_payload(
- self, output: str, api_type: AzureMLEndpointApiType
- ) -> str:
- response_json = json.loads(output)
- return response_json[0]["0"]
- """
- content_type: Optional[str] = "application/json"
- """The MIME type of the input data passed to the endpoint"""
-
- accepts: Optional[str] = "application/json"
- """The MIME type of the response data returned from the endpoint"""
-
- format_error_msg: str = (
- "Error while formatting response payload for chat model of type "
- " `{api_type}`. Are you using the right formatter for the deployed "
- " model and endpoint type?"
- )
-
- @staticmethod
- def escape_special_characters(prompt: str) -> str:
- """Escapes any special characters in `prompt`"""
- escape_map = {
- "\\": "\\\\",
- '"': '\\"',
- "\b": "\\b",
- "\f": "\\f",
- "\n": "\\n",
- "\r": "\\r",
- "\t": "\\t",
- }
-
- # Replace each occurrence of the specified characters with escaped versions
- for escape_sequence, escaped_sequence in escape_map.items():
- prompt = prompt.replace(escape_sequence, escaped_sequence)
-
- return prompt
-
- @property
- def supported_api_types(self) -> List[AzureMLEndpointApiType]:
- """Supported APIs for the given formatter. Azure ML supports
- deploying models using different hosting methods. Each method may have
- a different API structure."""
-
- return [AzureMLEndpointApiType.dedicated]
-
- def format_request_payload(
- self,
- prompt: str,
- model_kwargs: Dict,
- api_type: AzureMLEndpointApiType = AzureMLEndpointApiType.dedicated,
- ) -> Any:
- """Formats the request body according to the input schema of
- the model. Returns bytes or seekable file like object in the
- format specified in the content_type request header.
- """
- raise NotImplementedError()
-
- @abstractmethod
- def format_response_payload(
- self,
- output: bytes,
- api_type: AzureMLEndpointApiType = AzureMLEndpointApiType.dedicated,
- ) -> Generation:
- """Formats the response body according to the output
- schema of the model. Returns the data type that is
- received from the response.
- """
-
-
-class GPT2ContentFormatter(ContentFormatterBase):
- """Content handler for GPT2"""
-
- @property
- def supported_api_types(self) -> List[AzureMLEndpointApiType]:
- return [AzureMLEndpointApiType.dedicated]
-
- def format_request_payload( # type: ignore[override]
- self, prompt: str, model_kwargs: Dict, api_type: AzureMLEndpointApiType
- ) -> bytes:
- prompt = ContentFormatterBase.escape_special_characters(prompt)
- request_payload = json.dumps(
- {
- "inputs": {"input_string": [f'"{prompt}"']},
- "parameters": model_kwargs,
- }
- )
- return str.encode(request_payload)
-
- def format_response_payload( # type: ignore[override]
- self, output: bytes, api_type: AzureMLEndpointApiType
- ) -> Generation:
- try:
- choice = json.loads(output)[0]["0"]
- except (KeyError, IndexError, TypeError) as e:
- raise ValueError(self.format_error_msg.format(api_type=api_type)) from e
- return Generation(text=choice)
-
-
-class OSSContentFormatter(GPT2ContentFormatter):
- """Deprecated: Kept for backwards compatibility
-
- Content handler for LLMs from the OSS catalog."""
-
- content_formatter: Any = None
-
- def __init__(self) -> None:
- super().__init__()
- warnings.warn(
- """`OSSContentFormatter` will be deprecated in the future.
- Please use `GPT2ContentFormatter` instead.
- """
- )
-
-
-class HFContentFormatter(ContentFormatterBase):
- """Content handler for LLMs from the HuggingFace catalog."""
-
- @property
- def supported_api_types(self) -> List[AzureMLEndpointApiType]:
- return [AzureMLEndpointApiType.dedicated]
-
- def format_request_payload( # type: ignore[override]
- self, prompt: str, model_kwargs: Dict, api_type: AzureMLEndpointApiType
- ) -> bytes:
- ContentFormatterBase.escape_special_characters(prompt)
- request_payload = json.dumps(
- {
- "inputs": [f'"{prompt}"'],
- "parameters": model_kwargs,
- }
- )
- return str.encode(request_payload)
-
- def format_response_payload( # type: ignore[override]
- self, output: bytes, api_type: AzureMLEndpointApiType
- ) -> Generation:
- try:
- choice = json.loads(output)[0]["0"]["generated_text"]
- except (KeyError, IndexError, TypeError) as e:
- raise ValueError(self.format_error_msg.format(api_type=api_type)) from e
- return Generation(text=choice)
-
-
-class DollyContentFormatter(ContentFormatterBase):
- """Content handler for the Dolly-v2-12b model"""
-
- @property
- def supported_api_types(self) -> List[AzureMLEndpointApiType]:
- return [AzureMLEndpointApiType.dedicated]
-
- def format_request_payload( # type: ignore[override]
- self, prompt: str, model_kwargs: Dict, api_type: AzureMLEndpointApiType
- ) -> bytes:
- prompt = ContentFormatterBase.escape_special_characters(prompt)
- request_payload = json.dumps(
- {
- "input_data": {"input_string": [f'"{prompt}"']},
- "parameters": model_kwargs,
- }
- )
- return str.encode(request_payload)
-
- def format_response_payload( # type: ignore[override]
- self, output: bytes, api_type: AzureMLEndpointApiType
- ) -> Generation:
- try:
- choice = json.loads(output)[0]
- except (KeyError, IndexError, TypeError) as e:
- raise ValueError(self.format_error_msg.format(api_type=api_type)) from e
- return Generation(text=choice)
-
-
-class CustomOpenAIContentFormatter(ContentFormatterBase):
- """Content formatter for models that use the OpenAI like API scheme."""
-
- @property
- def supported_api_types(self) -> List[AzureMLEndpointApiType]:
- return [AzureMLEndpointApiType.dedicated, AzureMLEndpointApiType.serverless]
-
- def format_request_payload( # type: ignore[override]
- self, prompt: str, model_kwargs: Dict, api_type: AzureMLEndpointApiType
- ) -> bytes:
- """Formats the request according to the chosen api"""
- prompt = ContentFormatterBase.escape_special_characters(prompt)
- if api_type in [
- AzureMLEndpointApiType.dedicated,
- AzureMLEndpointApiType.realtime,
- ]:
- request_payload = json.dumps(
- {
- "input_data": {
- "input_string": [f'"{prompt}"'],
- "parameters": model_kwargs,
- }
- }
- )
- elif api_type == AzureMLEndpointApiType.serverless:
- request_payload = json.dumps({"prompt": prompt, **model_kwargs})
- else:
- raise ValueError(
- f"`api_type` {api_type} is not supported by this formatter"
- )
- return str.encode(request_payload)
-
- def format_response_payload( # type: ignore[override]
- self, output: bytes, api_type: AzureMLEndpointApiType
- ) -> Generation:
- """Formats response"""
- if api_type in [
- AzureMLEndpointApiType.dedicated,
- AzureMLEndpointApiType.realtime,
- ]:
- try:
- choice = json.loads(output)[0]["0"]
- except (KeyError, IndexError, TypeError) as e:
- raise ValueError(self.format_error_msg.format(api_type=api_type)) from e
- return Generation(text=choice)
- if api_type == AzureMLEndpointApiType.serverless:
- try:
- choice = json.loads(output)["choices"][0]
- if not isinstance(choice, dict):
- raise TypeError(
- "Endpoint response is not well formed for a chat "
- "model. Expected `dict` but `{type(choice)}` was "
- "received."
- )
- except (KeyError, IndexError, TypeError) as e:
- raise ValueError(self.format_error_msg.format(api_type=api_type)) from e
- return Generation(
- text=choice["text"].strip(),
- generation_info=dict(
- finish_reason=choice.get("finish_reason"),
- logprobs=choice.get("logprobs"),
- ),
- )
- raise ValueError(f"`api_type` {api_type} is not supported by this formatter")
-
-
-class LlamaContentFormatter(CustomOpenAIContentFormatter):
- """Deprecated: Kept for backwards compatibility
-
- Content formatter for Llama."""
-
- content_formatter: Any = None
-
- def __init__(self) -> None:
- super().__init__()
- warnings.warn(
- """`LlamaContentFormatter` will be deprecated in the future.
- Please use `CustomOpenAIContentFormatter` instead.
- """
- )
-
-
-class AzureMLBaseEndpoint(BaseModel):
- """Azure ML Online Endpoint models."""
-
- endpoint_url: str = ""
- """URL of pre-existing Endpoint. Should be passed to constructor or specified as
- env var `AZUREML_ENDPOINT_URL`."""
-
- endpoint_api_type: AzureMLEndpointApiType = AzureMLEndpointApiType.dedicated
- """Type of the endpoint being consumed. Possible values are `serverless` for
- pay-as-you-go and `dedicated` for dedicated endpoints. """
-
- endpoint_api_key: SecretStr = convert_to_secret_str("")
- """Authentication Key for Endpoint. Should be passed to constructor or specified as
- env var `AZUREML_ENDPOINT_API_KEY`."""
-
- deployment_name: str = ""
- """Deployment Name for Endpoint. NOT REQUIRED to call endpoint. Should be passed
- to constructor or specified as env var `AZUREML_DEPLOYMENT_NAME`."""
-
- timeout: int = DEFAULT_TIMEOUT
- """Request timeout for calls to the endpoint"""
-
- http_client: Any = None #: :meta private:
-
- max_retries: int = 1
-
- content_formatter: Any = None
- """The content formatter that provides an input and output
- transform function to handle formats between the LLM and
- the endpoint"""
-
- model_kwargs: Optional[dict] = None
- """Keyword arguments to pass to the model."""
-
- model_config = ConfigDict(protected_namespaces=())
-
- @model_validator(mode="before")
- @classmethod
- def validate_environ(cls, values: Dict) -> Any:
- values["endpoint_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(values, "endpoint_api_key", "AZUREML_ENDPOINT_API_KEY")
- )
- values["endpoint_url"] = get_from_dict_or_env(
- values, "endpoint_url", "AZUREML_ENDPOINT_URL"
- )
- values["deployment_name"] = get_from_dict_or_env(
- values, "deployment_name", "AZUREML_DEPLOYMENT_NAME", ""
- )
- values["endpoint_api_type"] = get_from_dict_or_env(
- values,
- "endpoint_api_type",
- "AZUREML_ENDPOINT_API_TYPE",
- AzureMLEndpointApiType.dedicated,
- )
- values["timeout"] = get_from_dict_or_env(
- values,
- "timeout",
- "AZUREML_TIMEOUT",
- str(DEFAULT_TIMEOUT),
- )
-
- return values
-
- @validator("content_formatter")
- def validate_content_formatter(
- cls, field_value: Any, values: Dict
- ) -> ContentFormatterBase:
- """Validate that content formatter is supported by endpoint type."""
- endpoint_api_type = values.get("endpoint_api_type")
- if endpoint_api_type not in field_value.supported_api_types:
- raise ValueError(
- f"Content formatter f{type(field_value)} is not supported by this "
- f"endpoint. Supported types are {field_value.supported_api_types} "
- f"but endpoint is {endpoint_api_type}."
- )
- return field_value
-
- @validator("endpoint_url")
- def validate_endpoint_url(cls, field_value: Any) -> str:
- """Validate that endpoint url is complete."""
- if field_value.endswith("/"):
- field_value = field_value[:-1]
- if field_value.endswith("inference.ml.azure.com"):
- raise ValueError(
- "`endpoint_url` should contain the full invocation URL including "
- "`/score` for `endpoint_api_type='dedicated'` or `/completions` "
- "or `/models/chat/completions` "
- "for `endpoint_api_type='serverless'`"
- )
- return field_value
-
- @validator("endpoint_api_type")
- def validate_endpoint_api_type(
- cls, field_value: Any, values: Dict
- ) -> AzureMLEndpointApiType:
- """Validate that endpoint api type is compatible with the URL format."""
- endpoint_url = values.get("endpoint_url")
- if (
- (
- field_value == AzureMLEndpointApiType.dedicated
- or field_value == AzureMLEndpointApiType.realtime
- )
- and not endpoint_url.endswith("/score") # type: ignore[union-attr]
- ):
- raise ValueError(
- "Endpoints of type `dedicated` should follow the format "
- "`https://..inference.ml.azure.com/score`."
- " If your endpoint URL ends with `/completions` or"
- "`/models/chat/completions`,"
- "use `endpoint_api_type='serverless'` instead."
- )
- if field_value == AzureMLEndpointApiType.serverless and not (
- endpoint_url.endswith("/completions") # type: ignore[union-attr]
- or endpoint_url.endswith("/models/chat/completions") # type: ignore[union-attr]
- ):
- raise ValueError(
- "Endpoints of type `serverless` should follow the format "
- "`https://..inference.ml.azure.com/completions`"
- " or `https://..inference.ml.azure.com/models/chat/completions`"
- )
-
- return field_value
-
- @validator("http_client", always=True)
- def validate_client(cls, field_value: Any, values: Dict) -> AzureMLEndpointClient:
- """Validate that api key and python package exists in environment."""
- endpoint_url = values.get("endpoint_url")
- endpoint_key = values.get("endpoint_api_key")
- deployment_name = values.get("deployment_name")
- timeout = values.get("timeout", DEFAULT_TIMEOUT)
-
- http_client = AzureMLEndpointClient(
- endpoint_url, # type: ignore[arg-type]
- endpoint_key.get_secret_value(), # type: ignore[union-attr]
- deployment_name, # type: ignore[arg-type]
- timeout,
- )
-
- return http_client
-
-
-class AzureMLOnlineEndpoint(BaseLLM, AzureMLBaseEndpoint):
- """Azure ML Online Endpoint models.
-
- Example:
- .. code-block:: python
- azure_llm = AzureMLOnlineEndpoint(
- endpoint_url="https://..inference.ml.azure.com/score",
- endpoint_api_type=AzureMLApiType.dedicated,
- endpoint_api_key="my-api-key",
- timeout=120,
- content_formatter=content_formatter,
- )
- """
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- **{"deployment_name": self.deployment_name},
- **{"model_kwargs": _model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "azureml_endpoint"
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Run the LLM on the given prompts.
-
- Args:
- prompts: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
- Returns:
- The string generated by the model.
- Example:
- .. code-block:: python
- response = azureml_model.invoke("Tell me a joke.")
- """
- _model_kwargs = self.model_kwargs or {}
- _model_kwargs.update(kwargs)
- if stop:
- _model_kwargs["stop"] = stop
- generations = []
-
- for prompt in prompts:
- request_payload = self.content_formatter.format_request_payload(
- prompt, _model_kwargs, self.endpoint_api_type
- )
- response_payload = self.http_client.call(
- body=request_payload, run_manager=run_manager
- )
- generated_text = self.content_formatter.format_response_payload(
- response_payload, self.endpoint_api_type
- )
- generations.append([generated_text])
-
- return LLMResult(generations=generations)
diff --git a/libs/community/langchain_community/llms/baichuan.py b/libs/community/langchain_community/llms/baichuan.py
deleted file mode 100644
index 4026e14b90..0000000000
--- a/libs/community/langchain_community/llms/baichuan.py
+++ /dev/null
@@ -1,95 +0,0 @@
-from __future__ import annotations
-
-import json
-import logging
-from typing import Any, Dict, List, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import Field, SecretStr
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-class BaichuanLLM(LLM):
- # TODO: Adding streaming support.
- """Baichuan large language models."""
-
- model: str = "Baichuan2-Turbo-192k"
- """
- Other models are available at https://platform.baichuan-ai.com/docs/api.
- """
- temperature: float = 0.3
- top_p: float = 0.95
- timeout: int = 60
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
-
- baichuan_api_host: Optional[str] = None
- baichuan_api_key: Optional[SecretStr] = None
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- values["baichuan_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(values, "baichuan_api_key", "BAICHUAN_API_KEY")
- )
- values["baichuan_api_host"] = get_from_dict_or_env(
- values,
- "baichuan_api_host",
- "BAICHUAN_API_HOST",
- default="https://api.baichuan-ai.com/v1/chat/completions",
- )
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- return {
- "model": self.model,
- "temperature": self.temperature,
- "top_p": self.top_p,
- **self.model_kwargs,
- }
-
- def _post(self, request: Any) -> Any:
- headers = {
- "Content-Type": "application/json",
- "Authorization": f"Bearer {self.baichuan_api_key.get_secret_value()}", # type: ignore[union-attr]
- }
- try:
- response = requests.post(
- self.baichuan_api_host, # type: ignore[arg-type]
- headers=headers,
- json=request,
- timeout=self.timeout,
- )
-
- if response.status_code == 200:
- parsed_json = json.loads(response.text)
- return parsed_json["choices"][0]["message"]["content"]
- else:
- response.raise_for_status()
- except Exception as e:
- raise ValueError(f"An error has occurred: {e}")
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- request = self._default_params
- request["messages"] = [{"role": "user", "content": prompt}]
- request.update(kwargs)
- text = self._post(request)
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
- return text
-
- @property
- def _llm_type(self) -> str:
- """Return type of chat_model."""
- return "baichuan-llm"
diff --git a/libs/community/langchain_community/llms/baidu_qianfan_endpoint.py b/libs/community/langchain_community/llms/baidu_qianfan_endpoint.py
deleted file mode 100644
index 5268b78034..0000000000
--- a/libs/community/langchain_community/llms/baidu_qianfan_endpoint.py
+++ /dev/null
@@ -1,315 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import (
- Any,
- AsyncIterator,
- Dict,
- Iterator,
- List,
- Optional,
-)
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import Field, SecretStr
-
-logger = logging.getLogger(__name__)
-
-
-class QianfanLLMEndpoint(LLM):
- """Baidu Qianfan completion model integration.
-
- Setup:
- Install ``qianfan`` and set environment variables ``QIANFAN_AK``, ``QIANFAN_SK``.
-
- .. code-block:: bash
-
- pip install qianfan
- export QIANFAN_AK="your-api-key"
- export QIANFAN_SK="your-secret_key"
-
- Key init args — completion params:
- model: str
- Name of Qianfan model to use.
- temperature: Optional[float]
- Sampling temperature.
- endpoint: Optional[str]
- Endpoint of the Qianfan LLM
- top_p: Optional[float]
- What probability mass to use.
-
- Key init args — client params:
- timeout: Optional[int]
- Timeout for requests.
- api_key: Optional[str]
- Qianfan API KEY. If not passed in will be read from env var QIANFAN_AK.
- secret_key: Optional[str]
- Qianfan SECRET KEY. If not passed in will be read from env var QIANFAN_SK.
-
- See full list of supported init args and their descriptions in the params section.
-
- Instantiate:
- .. code-block:: python
-
- from langchain_community.llms import QianfanLLMEndpoint
-
- llm = QianfanLLMEndpoint(
- model="ERNIE-3.5-8K",
- # api_key="...",
- # secret_key="...",
- # other params...
- )
-
- Invoke:
- .. code-block:: python
-
- input_text = "用50个字左右阐述,生命的意义在于"
- llm.invoke(input_text)
-
- .. code-block:: python
-
- '生命的意义在于体验、成长、爱与被爱、贡献与传承,以及对未知的勇敢探索与自我超越。'
-
- Stream:
- .. code-block:: python
-
- for chunk in llm.stream(input_text):
- print(chunk)
-
- .. code-block:: python
-
- 生命的意义 | 在于不断探索 | 与成长 | ,实现 | 自我价值,| 给予爱 | 并接受 | 爱, | 在经历 | 中感悟 | ,让 | 短暂的存在 | 绽放出无限 | 的光彩 | 与温暖 | 。
-
- .. code-block:: python
-
- stream = llm.stream(input_text)
- full = next(stream)
- for chunk in stream:
- full += chunk
- full
-
- .. code-block::
-
- '生命的意义在于探索、成长、爱与被爱、贡献价值、体验世界之美,以及在有限的时间里追求内心的平和与幸福。'
-
- Async:
- .. code-block:: python
-
- await llm.ainvoke(input_text)
-
- # stream:
- # async for chunk in llm.astream(input_text):
- # print(chunk)
-
- # batch:
- # await llm.abatch([input_text])
-
- .. code-block:: python
-
- '生命的意义在于探索、成长、爱与被爱、贡献社会,在有限的时间里追寻无限的可能,实现自我价值,让生活充满色彩与意义。'
-
- """ # noqa: E501
-
- init_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """init kwargs for qianfan client init, such as `query_per_second` which is
- associated with qianfan resource object to limit QPS"""
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """extra params for model invoke using with `do`."""
-
- client: Any = None
-
- qianfan_ak: Optional[SecretStr] = Field(default=None, alias="api_key")
- qianfan_sk: Optional[SecretStr] = Field(default=None, alias="secret_key")
-
- streaming: Optional[bool] = False
- """Whether to stream the results or not."""
-
- model: Optional[str] = Field(default=None)
- """Model name.
- you could get from https://cloud.baidu.com/doc/WENXINWORKSHOP/s/Nlks5zkzu
-
- preset models are mapping to an endpoint.
- `model` will be ignored if `endpoint` is set
-
- Default is set by `qianfan` SDK, not here
- """
-
- endpoint: Optional[str] = None
- """Endpoint of the Qianfan LLM, required if custom model used."""
-
- request_timeout: Optional[int] = Field(default=60, alias="timeout")
- """request timeout for chat http requests"""
-
- top_p: Optional[float] = 0.8
- temperature: Optional[float] = 0.95
- penalty_score: Optional[float] = 1
- """Model params, only supported in ERNIE-Bot and ERNIE-Bot-turbo.
- In the case of other model, passing these params will not affect the result.
- """
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- values["qianfan_ak"] = convert_to_secret_str(
- get_from_dict_or_env(
- values,
- ["qianfan_ak", "api_key"],
- "QIANFAN_AK",
- default="",
- )
- )
- values["qianfan_sk"] = convert_to_secret_str(
- get_from_dict_or_env(
- values,
- ["qianfan_sk", "secret_key"],
- "QIANFAN_SK",
- default="",
- )
- )
-
- params = {
- **values.get("init_kwargs", {}),
- "model": values["model"],
- }
- if values["qianfan_ak"].get_secret_value() != "":
- params["ak"] = values["qianfan_ak"].get_secret_value()
- if values["qianfan_sk"].get_secret_value() != "":
- params["sk"] = values["qianfan_sk"].get_secret_value()
- if values["endpoint"] is not None and values["endpoint"] != "":
- params["endpoint"] = values["endpoint"]
- try:
- import qianfan
-
- values["client"] = qianfan.Completion(**params)
- except ImportError:
- raise ImportError(
- "qianfan package not found, please install it with "
- "`pip install qianfan`"
- )
- return values
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- return {
- **{"endpoint": self.endpoint, "model": self.model},
- **super()._identifying_params,
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "baidu-qianfan-endpoint"
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling Qianfan API."""
- normal_params = {
- "model": self.model,
- "endpoint": self.endpoint,
- "stream": self.streaming,
- "request_timeout": self.request_timeout,
- "top_p": self.top_p,
- "temperature": self.temperature,
- "penalty_score": self.penalty_score,
- }
-
- return {**normal_params, **self.model_kwargs}
-
- def _convert_prompt_msg_params(
- self,
- prompt: str,
- **kwargs: Any,
- ) -> dict:
- if "streaming" in kwargs:
- kwargs["stream"] = kwargs.pop("streaming")
- return {
- **{"prompt": prompt, "model": self.model},
- **self._default_params,
- **kwargs,
- }
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to an qianfan models endpoint for each generation with a prompt.
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
- response = qianfan_model.invoke("Tell me a joke.")
- """
- if self.streaming:
- completion = ""
- for chunk in self._stream(prompt, stop, run_manager, **kwargs):
- completion += chunk.text
- return completion
- params = self._convert_prompt_msg_params(prompt, **kwargs)
- params["stop"] = stop
- response_payload = self.client.do(**params)
-
- return response_payload["result"]
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- if self.streaming:
- completion = ""
- async for chunk in self._astream(prompt, stop, run_manager, **kwargs):
- completion += chunk.text
- return completion
-
- params = self._convert_prompt_msg_params(prompt, **kwargs)
- params["stop"] = stop
- response_payload = await self.client.ado(**params)
-
- return response_payload["result"]
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- params = self._convert_prompt_msg_params(prompt, **{**kwargs, "stream": True})
- params["stop"] = stop
- for res in self.client.do(**params):
- if res:
- chunk = GenerationChunk(text=res["result"])
- if run_manager:
- run_manager.on_llm_new_token(chunk.text)
- yield chunk
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- params = self._convert_prompt_msg_params(prompt, **{**kwargs, "stream": True})
- params["stop"] = stop
- async for res in await self.client.ado(**params):
- if res:
- chunk = GenerationChunk(text=res["result"])
- if run_manager:
- await run_manager.on_llm_new_token(chunk.text)
- yield chunk
diff --git a/libs/community/langchain_community/llms/bananadev.py b/libs/community/langchain_community/llms/bananadev.py
deleted file mode 100644
index 9ed76e816b..0000000000
--- a/libs/community/langchain_community/llms/bananadev.py
+++ /dev/null
@@ -1,130 +0,0 @@
-import logging
-from typing import Any, Dict, List, Mapping, Optional, cast
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import (
- secret_from_env,
-)
-from pydantic import ConfigDict, Field, SecretStr, model_validator
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-class Banana(LLM):
- """Banana large language models.
-
- To use, you should have the ``banana-dev`` python package installed,
- and the environment variable ``BANANA_API_KEY`` set with your API key.
- This is the team API key available in the Banana dashboard.
-
- Any parameters that are valid to be passed to the call can be passed
- in, even if not explicitly saved on this class.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import Banana
- banana = Banana(model_key="", model_url_slug="")
- """
-
- model_key: str = ""
- """model key to use"""
-
- model_url_slug: str = ""
- """model endpoint to use"""
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `create` call not
- explicitly specified."""
-
- banana_api_key: Optional[SecretStr] = Field(
- default_factory=secret_from_env("BANANA_API_KEY", default=None)
- )
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = set(list(cls.model_fields.keys()))
- extra = values.get("model_kwargs", {})
- for field_name in list(values):
- if field_name not in all_required_field_names:
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- logger.warning(
- f"""{field_name} was transferred to model_kwargs.
- Please confirm that {field_name} is what you intended."""
- )
- extra[field_name] = values.pop(field_name)
- values["model_kwargs"] = extra
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- **{"model_key": self.model_key},
- **{"model_url_slug": self.model_url_slug},
- **{"model_kwargs": self.model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "bananadev"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call to Banana endpoint."""
- try:
- from banana_dev import Client
- except ImportError:
- raise ImportError(
- "Could not import banana-dev python package. "
- "Please install it with `pip install banana-dev`."
- )
- params = self.model_kwargs or {}
- params = {**params, **kwargs}
- api_key = cast(SecretStr, self.banana_api_key)
- model_key = self.model_key
- model_url_slug = self.model_url_slug
- model_inputs = {
- # a json specific to your model.
- "prompt": prompt,
- **params,
- }
- model = Client(
- # Found in main dashboard
- api_key=api_key.get_secret_value(),
- # Both found in model details page
- model_key=model_key,
- url=f"https://{model_url_slug}.run.banana.dev",
- )
- response, meta = model.call("/", model_inputs)
- try:
- text = response["outputs"]
- except (KeyError, TypeError):
- raise ValueError(
- "Response should be of schema: {'outputs': 'text'}."
- "\nTo fix this:"
- "\n- fork the source repo of the Banana model"
- "\n- modify app.py to return the above schema"
- "\n- deploy that as a custom repo"
- )
- if stop is not None:
- # I believe this is required since the stop tokens
- # are not enforced by the model parameters
- text = enforce_stop_tokens(text, stop)
- return text
diff --git a/libs/community/langchain_community/llms/baseten.py b/libs/community/langchain_community/llms/baseten.py
deleted file mode 100644
index 5b7ce87eb8..0000000000
--- a/libs/community/langchain_community/llms/baseten.py
+++ /dev/null
@@ -1,94 +0,0 @@
-import logging
-import os
-from typing import Any, Dict, List, Mapping, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from pydantic import Field
-
-logger = logging.getLogger(__name__)
-
-
-class Baseten(LLM):
- """Baseten model
-
- This module allows using LLMs hosted on Baseten.
-
- The LLM deployed on Baseten must have the following properties:
-
- * Must accept input as a dictionary with the key "prompt"
- * May accept other input in the dictionary passed through with kwargs
- * Must return a string with the model output
-
- To use this module, you must:
-
- * Export your Baseten API key as the environment variable `BASETEN_API_KEY`
- * Get the model ID for your model from your Baseten dashboard
- * Identify the model deployment ("production" for all model library models)
-
- These code samples use
- [Mistral 7B Instruct](https://app.baseten.co/explore/mistral_7b_instruct)
- from Baseten's model library.
-
- Examples:
- .. code-block:: python
-
- from langchain_community.llms import Baseten
- # Production deployment
- mistral = Baseten(model="MODEL_ID", deployment="production")
- mistral("What is the Mistral wind?")
-
- .. code-block:: python
-
- from langchain_community.llms import Baseten
- # Development deployment
- mistral = Baseten(model="MODEL_ID", deployment="development")
- mistral("What is the Mistral wind?")
-
- .. code-block:: python
-
- from langchain_community.llms import Baseten
- # Other published deployment
- mistral = Baseten(model="MODEL_ID", deployment="DEPLOYMENT_ID")
- mistral("What is the Mistral wind?")
- """
-
- model: str
- deployment: str
- input: Dict[str, Any] = Field(default_factory=dict)
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- **{"model_kwargs": self.model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of model."""
- return "baseten"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- baseten_api_key = os.environ["BASETEN_API_KEY"]
- model_id = self.model
- if self.deployment == "production":
- model_url = f"https://model-{model_id}.api.baseten.co/production/predict"
- elif self.deployment == "development":
- model_url = f"https://model-{model_id}.api.baseten.co/development/predict"
- else: # try specific deployment ID
- model_url = f"https://model-{model_id}.api.baseten.co/deployment/{self.deployment}/predict"
- response = requests.post(
- model_url,
- headers={"Authorization": f"Api-Key {baseten_api_key}"},
- json={"prompt": prompt, **kwargs},
- )
- return response.json()
diff --git a/libs/community/langchain_community/llms/beam.py b/libs/community/langchain_community/llms/beam.py
deleted file mode 100644
index 7d3b6e65ec..0000000000
--- a/libs/community/langchain_community/llms/beam.py
+++ /dev/null
@@ -1,273 +0,0 @@
-import base64
-import json
-import logging
-import subprocess
-import textwrap
-import time
-from typing import Any, Dict, List, Mapping, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from langchain_core.utils.pydantic import get_fields
-from pydantic import ConfigDict, Field, model_validator
-
-logger = logging.getLogger(__name__)
-
-DEFAULT_NUM_TRIES = 10
-DEFAULT_SLEEP_TIME = 4
-
-
-class Beam(LLM):
- """Beam API for gpt2 large language model.
-
- To use, you should have the ``beam-sdk`` python package installed,
- and the environment variable ``BEAM_CLIENT_ID`` set with your client id
- and ``BEAM_CLIENT_SECRET`` set with your client secret. Information on how
- to get this is available here: https://docs.beam.cloud/account/api-keys.
-
- The wrapper can then be called as follows, where the name, cpu, memory, gpu,
- python version, and python packages can be updated accordingly. Once deployed,
- the instance can be called.
-
- Example:
- .. code-block:: python
-
- llm = Beam(model_name="gpt2",
- name="langchain-gpt2",
- cpu=8,
- memory="32Gi",
- gpu="A10G",
- python_version="python3.8",
- python_packages=[
- "diffusers[torch]>=0.10",
- "transformers",
- "torch",
- "pillow",
- "accelerate",
- "safetensors",
- "xformers",],
- max_length=50)
- llm._deploy()
- call_result = llm._call(input)
-
- """
-
- model_name: str = ""
- name: str = ""
- cpu: str = ""
- memory: str = ""
- gpu: str = ""
- python_version: str = ""
- python_packages: List[str] = []
- max_length: str = ""
- url: str = ""
- """model endpoint to use"""
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `create` call not
- explicitly specified."""
-
- beam_client_id: str = ""
- beam_client_secret: str = ""
- app_id: Optional[str] = None
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = {field.alias for field in get_fields(cls).values()}
-
- extra = values.get("model_kwargs", {})
- for field_name in list(values):
- if field_name not in all_required_field_names:
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- logger.warning(
- f"""{field_name} was transferred to model_kwargs.
- Please confirm that {field_name} is what you intended."""
- )
- extra[field_name] = values.pop(field_name)
- values["model_kwargs"] = extra
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- beam_client_id = get_from_dict_or_env(
- values, "beam_client_id", "BEAM_CLIENT_ID"
- )
- beam_client_secret = get_from_dict_or_env(
- values, "beam_client_secret", "BEAM_CLIENT_SECRET"
- )
- values["beam_client_id"] = beam_client_id
- values["beam_client_secret"] = beam_client_secret
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- "model_name": self.model_name,
- "name": self.name,
- "cpu": self.cpu,
- "memory": self.memory,
- "gpu": self.gpu,
- "python_version": self.python_version,
- "python_packages": self.python_packages,
- "max_length": self.max_length,
- "model_kwargs": self.model_kwargs,
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "beam"
-
- def app_creation(self) -> None:
- """Creates a Python file which will contain your Beam app definition."""
- script = textwrap.dedent(
- """\
- import beam
-
- # The environment your code will run on
- app = beam.App(
- name="{name}",
- cpu={cpu},
- memory="{memory}",
- gpu="{gpu}",
- python_version="{python_version}",
- python_packages={python_packages},
- )
-
- app.Trigger.RestAPI(
- inputs={{"prompt": beam.Types.String(), "max_length": beam.Types.String()}},
- outputs={{"text": beam.Types.String()}},
- handler="run.py:beam_langchain",
- )
-
- """
- )
-
- script_name = "app.py"
- with open(script_name, "w") as file:
- file.write(
- script.format(
- name=self.name,
- cpu=self.cpu,
- memory=self.memory,
- gpu=self.gpu,
- python_version=self.python_version,
- python_packages=self.python_packages,
- )
- )
-
- def run_creation(self) -> None:
- """Creates a Python file which will be deployed on beam."""
- script = textwrap.dedent(
- """
- import os
- import transformers
- from transformers import GPT2LMHeadModel, GPT2Tokenizer
-
- model_name = "{model_name}"
-
- def beam_langchain(**inputs):
- prompt = inputs["prompt"]
- length = inputs["max_length"]
-
- tokenizer = GPT2Tokenizer.from_pretrained(model_name)
- model = GPT2LMHeadModel.from_pretrained(model_name)
- encodedPrompt = tokenizer.encode(prompt, return_tensors='pt')
- outputs = model.generate(encodedPrompt, max_length=int(length),
- do_sample=True, pad_token_id=tokenizer.eos_token_id)
- output = tokenizer.decode(outputs[0], skip_special_tokens=True)
-
- print(output) # noqa: T201
- return {{"text": output}}
-
- """
- )
-
- script_name = "run.py"
- with open(script_name, "w") as file:
- file.write(script.format(model_name=self.model_name))
-
- def _deploy(self) -> str:
- """Call to Beam."""
- try:
- import beam
-
- if beam.__path__ == "":
- raise ImportError
- except ImportError:
- raise ImportError(
- "Could not import beam python package. "
- "Please install it with `curl "
- "https://raw.githubusercontent.com/slai-labs"
- "/get-beam/main/get-beam.sh -sSfL | sh`."
- )
- self.app_creation()
- self.run_creation()
-
- process = subprocess.run(
- "beam deploy app.py", shell=True, capture_output=True, text=True
- )
-
- if process.returncode == 0:
- output = process.stdout
- logger.info(output)
- lines = output.split("\n")
-
- for line in lines:
- if line.startswith(" i Send requests to: https://apps.beam.cloud/"):
- self.app_id = line.split("/")[-1]
- self.url = line.split(":")[1].strip()
- return self.app_id
-
- raise ValueError(
- f"""Failed to retrieve the appID from the deployment output.
- Deployment output: {output}"""
- )
- else:
- raise ValueError(f"Deployment failed. Error: {process.stderr}")
-
- @property
- def authorization(self) -> str:
- if self.beam_client_id:
- credential_str = self.beam_client_id + ":" + self.beam_client_secret
- else:
- credential_str = self.beam_client_secret
- return base64.b64encode(credential_str.encode()).decode()
-
- def _call(
- self,
- prompt: str,
- stop: Optional[list] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call to Beam."""
- url = "https://apps.beam.cloud/" + self.app_id if self.app_id else self.url
- payload = {"prompt": prompt, "max_length": self.max_length}
- payload.update(kwargs)
- headers = {
- "Accept": "*/*",
- "Accept-Encoding": "gzip, deflate",
- "Authorization": "Basic " + self.authorization,
- "Connection": "keep-alive",
- "Content-Type": "application/json",
- }
-
- for _ in range(DEFAULT_NUM_TRIES):
- request = requests.post(url, headers=headers, data=json.dumps(payload))
- if request.status_code == 200:
- return request.json()["text"]
- time.sleep(DEFAULT_SLEEP_TIME)
- logger.warning("Unable to successfully call model.")
- return ""
diff --git a/libs/community/langchain_community/llms/bedrock.py b/libs/community/langchain_community/llms/bedrock.py
deleted file mode 100644
index e1b6570cf3..0000000000
--- a/libs/community/langchain_community/llms/bedrock.py
+++ /dev/null
@@ -1,917 +0,0 @@
-import asyncio
-import json
-import warnings
-from abc import ABC
-from typing import (
- Any,
- AsyncGenerator,
- AsyncIterator,
- Dict,
- Iterator,
- List,
- Mapping,
- Optional,
- Tuple,
-)
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict, Field
-
-from langchain_community.llms.utils import enforce_stop_tokens
-from langchain_community.utilities.anthropic import (
- get_num_tokens_anthropic,
- get_token_ids_anthropic,
-)
-
-AMAZON_BEDROCK_TRACE_KEY = "amazon-bedrock-trace"
-GUARDRAILS_BODY_KEY = "amazon-bedrock-guardrailAssessment"
-HUMAN_PROMPT = "\n\nHuman:"
-ASSISTANT_PROMPT = "\n\nAssistant:"
-ALTERNATION_ERROR = (
- "Error: Prompt must alternate between '\n\nHuman:' and '\n\nAssistant:'."
-)
-
-
-def _add_newlines_before_ha(input_text: str) -> str:
- new_text = input_text
- for word in ["Human:", "Assistant:"]:
- new_text = new_text.replace(word, "\n\n" + word)
- for i in range(2):
- new_text = new_text.replace("\n\n\n" + word, "\n\n" + word)
- return new_text
-
-
-def _human_assistant_format(input_text: str) -> str:
- if input_text.count("Human:") == 0 or (
- input_text.find("Human:") > input_text.find("Assistant:")
- and "Assistant:" in input_text
- ):
- input_text = HUMAN_PROMPT + " " + input_text # SILENT CORRECTION
- if input_text.count("Assistant:") == 0:
- input_text = input_text + ASSISTANT_PROMPT # SILENT CORRECTION
- if input_text[: len("Human:")] == "Human:":
- input_text = "\n\n" + input_text
- input_text = _add_newlines_before_ha(input_text)
- count = 0
- # track alternation
- for i in range(len(input_text)):
- if input_text[i : i + len(HUMAN_PROMPT)] == HUMAN_PROMPT:
- if count % 2 == 0:
- count += 1
- else:
- warnings.warn(ALTERNATION_ERROR + f" Received {input_text}")
- if input_text[i : i + len(ASSISTANT_PROMPT)] == ASSISTANT_PROMPT:
- if count % 2 == 1:
- count += 1
- else:
- warnings.warn(ALTERNATION_ERROR + f" Received {input_text}")
-
- if count % 2 == 1: # Only saw Human, no Assistant
- input_text = input_text + ASSISTANT_PROMPT # SILENT CORRECTION
-
- return input_text
-
-
-def _stream_response_to_generation_chunk(
- stream_response: Dict[str, Any],
-) -> GenerationChunk:
- """Convert a stream response to a generation chunk."""
- if not stream_response["delta"]:
- return GenerationChunk(text="")
- return GenerationChunk(
- text=stream_response["delta"]["text"],
- generation_info=dict(
- finish_reason=stream_response.get("stop_reason", None),
- ),
- )
-
-
-class LLMInputOutputAdapter:
- """Adapter class to prepare the inputs from Langchain to a format
- that LLM model expects.
-
- It also provides helper function to extract
- the generated text from the model response."""
-
- provider_to_output_key_map = {
- "anthropic": "completion",
- "amazon": "outputText",
- "cohere": "text",
- "meta": "generation",
- "mistral": "outputs",
- }
-
- @classmethod
- def prepare_input(
- cls,
- provider: str,
- model_kwargs: Dict[str, Any],
- prompt: Optional[str] = None,
- system: Optional[str] = None,
- messages: Optional[List[Dict]] = None,
- ) -> Dict[str, Any]:
- input_body = {**model_kwargs}
- if provider == "anthropic":
- if messages:
- input_body["anthropic_version"] = "bedrock-2023-05-31"
- input_body["messages"] = messages
- if system:
- input_body["system"] = system
- if "max_tokens" not in input_body:
- input_body["max_tokens"] = 1024
- if prompt:
- input_body["prompt"] = _human_assistant_format(prompt)
- if "max_tokens_to_sample" not in input_body:
- input_body["max_tokens_to_sample"] = 1024
- elif provider in ("ai21", "cohere", "meta", "mistral"):
- input_body["prompt"] = prompt
- elif provider == "amazon":
- input_body = dict()
- input_body["inputText"] = prompt
- input_body["textGenerationConfig"] = {**model_kwargs}
- else:
- input_body["inputText"] = prompt
-
- return input_body
-
- @classmethod
- def prepare_output(cls, provider: str, response: Any) -> dict:
- text = ""
- if provider == "anthropic":
- response_body = json.loads(response.get("body").read().decode())
- if "completion" in response_body:
- text = response_body.get("completion")
- elif "content" in response_body:
- content = response_body.get("content")
- text = content[0].get("text")
- else:
- response_body = json.loads(response.get("body").read())
-
- if provider == "ai21":
- text = response_body.get("completions")[0].get("data").get("text")
- elif provider == "cohere":
- text = response_body.get("generations")[0].get("text")
- elif provider == "meta":
- text = response_body.get("generation")
- elif provider == "mistral":
- text = response_body.get("outputs")[0].get("text")
- else:
- text = response_body.get("results")[0].get("outputText")
-
- headers = response.get("ResponseMetadata", {}).get("HTTPHeaders", {})
- prompt_tokens = int(headers.get("x-amzn-bedrock-input-token-count", 0))
- completion_tokens = int(headers.get("x-amzn-bedrock-output-token-count", 0))
- return {
- "text": text,
- "body": response_body,
- "usage": {
- "prompt_tokens": prompt_tokens,
- "completion_tokens": completion_tokens,
- "total_tokens": prompt_tokens + completion_tokens,
- },
- }
-
- @classmethod
- def prepare_output_stream(
- cls,
- provider: str,
- response: Any,
- stop: Optional[List[str]] = None,
- messages_api: bool = False,
- ) -> Iterator[GenerationChunk]:
- stream = response.get("body")
-
- if not stream:
- return
-
- if messages_api:
- output_key = "message"
- else:
- output_key = cls.provider_to_output_key_map.get(provider, "")
-
- if not output_key:
- raise ValueError(
- f"Unknown streaming response output key for provider: {provider}"
- )
-
- for event in stream:
- chunk = event.get("chunk")
- if not chunk:
- continue
-
- chunk_obj = json.loads(chunk.get("bytes").decode())
-
- if provider == "cohere" and (
- chunk_obj["is_finished"] or chunk_obj[output_key] == ""
- ):
- return
-
- elif (
- provider == "mistral"
- and chunk_obj.get(output_key, [{}])[0].get("stop_reason", "") == "stop"
- ):
- return
-
- elif messages_api and (chunk_obj.get("type") == "content_block_stop"):
- return
-
- if messages_api and chunk_obj.get("type") in (
- "message_start",
- "content_block_start",
- "content_block_delta",
- ):
- if chunk_obj.get("type") == "content_block_delta":
- chk = _stream_response_to_generation_chunk(chunk_obj)
- yield chk
- else:
- continue
- else:
- # chunk obj format varies with provider
- yield GenerationChunk(
- text=(
- chunk_obj[output_key]
- if provider != "mistral"
- else chunk_obj[output_key][0]["text"]
- ),
- generation_info={
- GUARDRAILS_BODY_KEY: (
- chunk_obj.get(GUARDRAILS_BODY_KEY)
- if GUARDRAILS_BODY_KEY in chunk_obj
- else None
- ),
- },
- )
-
- @classmethod
- async def aprepare_output_stream(
- cls, provider: str, response: Any, stop: Optional[List[str]] = None
- ) -> AsyncIterator[GenerationChunk]:
- stream = response.get("body")
-
- if not stream:
- return
-
- output_key = cls.provider_to_output_key_map.get(provider, None)
-
- if not output_key:
- raise ValueError(
- f"Unknown streaming response output key for provider: {provider}"
- )
-
- for event in stream:
- chunk = event.get("chunk")
- if not chunk:
- continue
-
- chunk_obj = json.loads(chunk.get("bytes").decode())
-
- if provider == "cohere" and (
- chunk_obj["is_finished"] or chunk_obj[output_key] == ""
- ):
- return
-
- if (
- provider == "mistral"
- and chunk_obj.get(output_key, [{}])[0].get("stop_reason", "") == "stop"
- ):
- return
-
- yield GenerationChunk(
- text=(
- chunk_obj[output_key]
- if provider != "mistral"
- else chunk_obj[output_key][0]["text"]
- )
- )
-
-
-class BedrockBase(BaseModel, ABC):
- """Base class for Bedrock models."""
-
- model_config = ConfigDict(protected_namespaces=())
-
- client: Any = Field(exclude=True) #: :meta private:
-
- region_name: Optional[str] = None
- """The aws region e.g., `us-west-2`. Fallsback to AWS_DEFAULT_REGION env variable
- or region specified in ~/.aws/config in case it is not provided here.
- """
-
- credentials_profile_name: Optional[str] = Field(default=None, exclude=True)
- """The name of the profile in the ~/.aws/credentials or ~/.aws/config files, which
- has either access keys or role information specified.
- If not specified, the default credential profile or, if on an EC2 instance,
- credentials from IMDS will be used.
- See: https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html
- """
-
- config: Any = None
- """An optional botocore.config.Config instance to pass to the client."""
-
- provider: Optional[str] = None
- """The model provider, e.g., amazon, cohere, ai21, etc. When not supplied, provider
- is extracted from the first part of the model_id e.g. 'amazon' in
- 'amazon.titan-text-express-v1'. This value should be provided for model ids that do
- not have the provider in them, e.g., custom and provisioned models that have an ARN
- associated with them."""
-
- model_id: str
- """Id of the model to call, e.g., amazon.titan-text-express-v1, this is
- equivalent to the modelId property in the list-foundation-models api. For custom and
- provisioned models, an ARN value is expected."""
-
- model_kwargs: Optional[Dict] = None
- """Keyword arguments to pass to the model."""
-
- endpoint_url: Optional[str] = None
- """Needed if you don't want to default to us-east-1 endpoint"""
-
- streaming: bool = False
- """Whether to stream the results."""
-
- provider_stop_sequence_key_name_map: Mapping[str, str] = {
- "anthropic": "stop_sequences",
- "amazon": "stopSequences",
- "ai21": "stop_sequences",
- "cohere": "stop_sequences",
- "mistral": "stop",
- }
-
- guardrails: Optional[Mapping[str, Any]] = {
- "id": None,
- "version": None,
- "trace": False,
- }
- """
- An optional dictionary to configure guardrails for Bedrock.
-
- This field 'guardrails' consists of two keys: 'id' and 'version',
- which should be strings, but are initialized to None. It's used to
- determine if specific guardrails are enabled and properly set.
-
- Type:
- Optional[Mapping[str, str]]: A mapping with 'id' and 'version' keys.
-
- Example:
- llm = Bedrock(model_id="", client=,
- model_kwargs={},
- guardrails={
- "id": "",
- "version": ""})
-
- To enable tracing for guardrails, set the 'trace' key to True and pass a callback handler to the
- 'run_manager' parameter of the 'generate', '_call' methods.
-
- Example:
- llm = Bedrock(model_id="", client=,
- model_kwargs={},
- guardrails={
- "id": "",
- "version": "",
- "trace": True},
- callbacks=[BedrockAsyncCallbackHandler()])
-
- [https://python.langchain.com/docs/modules/callbacks/] for more information on callback handlers.
-
- class BedrockAsyncCallbackHandler(AsyncCallbackHandler):
- async def on_llm_error(
- self,
- error: BaseException,
- **kwargs: Any,
- ) -> Any:
- reason = kwargs.get("reason")
- if reason == "GUARDRAIL_INTERVENED":
- ...Logic to handle guardrail intervention...
- """ # noqa: E501
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that AWS credentials to and python package exists in environment."""
-
- # Skip creating new client if passed in constructor
- if values.get("client") is not None:
- return values
-
- try:
- import boto3
-
- if values["credentials_profile_name"] is not None:
- session = boto3.Session(profile_name=values["credentials_profile_name"])
- else:
- # use default credentials
- session = boto3.Session()
-
- values["region_name"] = get_from_dict_or_env(
- values,
- "region_name",
- "AWS_DEFAULT_REGION",
- default=session.region_name,
- )
-
- client_params = {}
- if values["region_name"]:
- client_params["region_name"] = values["region_name"]
- if values["endpoint_url"]:
- client_params["endpoint_url"] = values["endpoint_url"]
- if values["config"]:
- client_params["config"] = values["config"]
-
- values["client"] = session.client("bedrock-runtime", **client_params)
-
- except ImportError:
- raise ImportError(
- "Could not import boto3 python package. "
- "Please install it with `pip install boto3`."
- )
- except ValueError as e:
- raise ValueError(f"Error raised by bedrock service: {e}")
- except Exception as e:
- raise ValueError(
- "Could not load credentials to authenticate with AWS client. "
- "Please check that credentials in the specified "
- f"profile name are valid. Bedrock error: {e}"
- ) from e
-
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- **{"model_kwargs": _model_kwargs},
- }
-
- def _get_provider(self) -> str:
- if self.provider:
- return self.provider
- if self.model_id.startswith("arn"):
- raise ValueError(
- "Model provider should be supplied when passing a model ARN as model_id"
- )
-
- return self.model_id.split(".")[0]
-
- @property
- def _model_is_anthropic(self) -> bool:
- return self._get_provider() == "anthropic"
-
- @property
- def _guardrails_enabled(self) -> bool:
- """
- Determines if guardrails are enabled and correctly configured.
- Checks if 'guardrails' is a dictionary with non-empty 'id' and 'version' keys.
- Checks if 'guardrails.trace' is true.
-
- Returns:
- bool: True if guardrails are correctly configured, False otherwise.
- Raises:
- TypeError: If 'guardrails' lacks 'id' or 'version' keys.
- """
- try:
- return (
- isinstance(self.guardrails, dict)
- and bool(self.guardrails["id"])
- and bool(self.guardrails["version"])
- )
-
- except KeyError as e:
- raise TypeError(
- "Guardrails must be a dictionary with 'id' and 'version' keys."
- ) from e
-
- def _get_guardrails_canonical(self) -> Dict[str, Any]:
- """
- The canonical way to pass in guardrails to the bedrock service
- adheres to the following format:
-
- "amazon-bedrock-guardrailDetails": {
- "guardrailId": "string",
- "guardrailVersion": "string"
- }
- """
- return {
- "amazon-bedrock-guardrailDetails": {
- "guardrailId": self.guardrails.get("id"), # type: ignore[union-attr]
- "guardrailVersion": self.guardrails.get("version"), # type: ignore[union-attr]
- }
- }
-
- def _prepare_input_and_invoke(
- self,
- prompt: Optional[str] = None,
- system: Optional[str] = None,
- messages: Optional[List[Dict]] = None,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Tuple[str, Dict[str, Any]]:
- _model_kwargs = self.model_kwargs or {}
-
- provider = self._get_provider()
- params = {**_model_kwargs, **kwargs}
- if self._guardrails_enabled:
- params.update(self._get_guardrails_canonical())
- input_body = LLMInputOutputAdapter.prepare_input(
- provider=provider,
- model_kwargs=params,
- prompt=prompt,
- system=system,
- messages=messages,
- )
- body = json.dumps(input_body)
- accept = "application/json"
- contentType = "application/json"
-
- request_options = {
- "body": body,
- "modelId": self.model_id,
- "accept": accept,
- "contentType": contentType,
- }
-
- if self._guardrails_enabled:
- request_options["guardrail"] = "ENABLED"
- if self.guardrails.get("trace"): # type: ignore[union-attr]
- request_options["trace"] = "ENABLED"
-
- try:
- response = self.client.invoke_model(**request_options)
-
- text, body, usage_info = LLMInputOutputAdapter.prepare_output(
- provider, response
- ).values()
-
- except Exception as e:
- raise ValueError(f"Error raised by bedrock service: {e}")
-
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
-
- # Verify and raise a callback error if any intervention occurs or a signal is
- # sent from a Bedrock service,
- # such as when guardrails are triggered.
- services_trace = self._get_bedrock_services_signal(body) # type: ignore[arg-type]
-
- if services_trace.get("signal") and run_manager is not None:
- run_manager.on_llm_error(
- Exception(
- f"Error raised by bedrock service: {services_trace.get('reason')}"
- ),
- **services_trace,
- )
-
- return text, usage_info
-
- def _get_bedrock_services_signal(self, body: dict) -> dict:
- """
- This function checks the response body for an interrupt flag or message that indicates
- whether any of the Bedrock services have intervened in the processing flow. It is
- primarily used to identify modifications or interruptions imposed by these services
- during the request-response cycle with a Large Language Model (LLM).
- """ # noqa: E501
-
- if (
- self._guardrails_enabled
- and self.guardrails.get("trace") # type: ignore[union-attr]
- and self._is_guardrails_intervention(body)
- ):
- return {
- "signal": True,
- "reason": "GUARDRAIL_INTERVENED",
- "trace": body.get(AMAZON_BEDROCK_TRACE_KEY),
- }
-
- return {
- "signal": False,
- "reason": None,
- "trace": None,
- }
-
- def _is_guardrails_intervention(self, body: dict) -> bool:
- return body.get(GUARDRAILS_BODY_KEY) == "GUARDRAIL_INTERVENED"
-
- def _prepare_input_and_invoke_stream(
- self,
- prompt: Optional[str] = None,
- system: Optional[str] = None,
- messages: Optional[List[Dict]] = None,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- _model_kwargs = self.model_kwargs or {}
- provider = self._get_provider()
-
- if stop:
- if provider not in self.provider_stop_sequence_key_name_map:
- raise ValueError(
- f"Stop sequence key name for {provider} is not supported."
- )
-
- # stop sequence from _generate() overrides
- # stop sequences in the class attribute
- _model_kwargs[self.provider_stop_sequence_key_name_map.get(provider)] = stop
-
- if provider == "cohere":
- _model_kwargs["stream"] = True
-
- params = {**_model_kwargs, **kwargs}
-
- if self._guardrails_enabled:
- params.update(self._get_guardrails_canonical())
-
- input_body = LLMInputOutputAdapter.prepare_input(
- provider=provider,
- prompt=prompt,
- system=system,
- messages=messages,
- model_kwargs=params,
- )
- body = json.dumps(input_body)
-
- request_options = {
- "body": body,
- "modelId": self.model_id,
- "accept": "application/json",
- "contentType": "application/json",
- }
-
- if self._guardrails_enabled:
- request_options["guardrail"] = "ENABLED"
- if self.guardrails.get("trace"): # type: ignore[union-attr]
- request_options["trace"] = "ENABLED"
-
- try:
- response = self.client.invoke_model_with_response_stream(**request_options)
-
- except Exception as e:
- raise ValueError(f"Error raised by bedrock service: {e}")
-
- for chunk in LLMInputOutputAdapter.prepare_output_stream(
- provider, response, stop, True if messages else False
- ):
- # verify and raise callback error if any middleware intervened
- self._get_bedrock_services_signal(chunk.generation_info) # type: ignore[arg-type]
-
- if run_manager is not None:
- run_manager.on_llm_new_token(chunk.text, chunk=chunk)
- yield chunk
-
- async def _aprepare_input_and_invoke_stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- _model_kwargs = self.model_kwargs or {}
- provider = self._get_provider()
-
- if stop:
- if provider not in self.provider_stop_sequence_key_name_map:
- raise ValueError(
- f"Stop sequence key name for {provider} is not supported."
- )
- _model_kwargs[self.provider_stop_sequence_key_name_map.get(provider)] = stop
-
- if provider == "cohere":
- _model_kwargs["stream"] = True
-
- params = {**_model_kwargs, **kwargs}
- input_body = LLMInputOutputAdapter.prepare_input(
- provider=provider, prompt=prompt, model_kwargs=params
- )
- body = json.dumps(input_body)
-
- response = await asyncio.get_running_loop().run_in_executor(
- None,
- lambda: self.client.invoke_model_with_response_stream(
- body=body,
- modelId=self.model_id,
- accept="application/json",
- contentType="application/json",
- ),
- )
-
- async for chunk in LLMInputOutputAdapter.aprepare_output_stream(
- provider, response, stop
- ):
- if run_manager is not None and asyncio.iscoroutinefunction(
- run_manager.on_llm_new_token
- ):
- await run_manager.on_llm_new_token(chunk.text, chunk=chunk)
- elif run_manager is not None:
- run_manager.on_llm_new_token(chunk.text, chunk=chunk) # type: ignore[unused-coroutine]
- yield chunk
-
-
-@deprecated(
- since="0.0.34", removal="1.0", alternative_import="langchain_aws.BedrockLLM"
-)
-class Bedrock(LLM, BedrockBase):
- """Bedrock models.
-
- To authenticate, the AWS client uses the following methods to
- automatically load credentials:
- https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html
-
- If a specific credential profile should be used, you must pass
- the name of the profile from the ~/.aws/credentials file that is to be used.
-
- Make sure the credentials / roles used have the required policies to
- access the Bedrock service.
- """
-
- """
- Example:
- .. code-block:: python
-
- from bedrock_langchain.bedrock_llm import BedrockLLM
-
- llm = BedrockLLM(
- credentials_profile_name="default",
- model_id="amazon.titan-text-express-v1",
- streaming=True
- )
-
- """
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- model_id = values["model_id"]
- if model_id.startswith("anthropic.claude-3"):
- raise ValueError(
- "Claude v3 models are not supported by this LLM."
- "Please use `from langchain_community.chat_models import BedrockChat` "
- "instead."
- )
- return super().validate_environment(values)
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "amazon_bedrock"
-
- @classmethod
- def is_lc_serializable(cls) -> bool:
- """Return whether this model can be serialized by Langchain."""
- return True
-
- @classmethod
- def get_lc_namespace(cls) -> List[str]:
- """Get the namespace of the langchain object."""
- return ["langchain", "llms", "bedrock"]
-
- @property
- def lc_attributes(self) -> Dict[str, Any]:
- attributes: Dict[str, Any] = {}
-
- if self.region_name:
- attributes["region_name"] = self.region_name
-
- return attributes
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- """Call out to Bedrock service with streaming.
-
- Args:
- prompt (str): The prompt to pass into the model
- stop (Optional[List[str]], optional): Stop sequences. These will
- override any stop sequences in the `model_kwargs` attribute.
- Defaults to None.
- run_manager (Optional[CallbackManagerForLLMRun], optional): Callback
- run managers used to process the output. Defaults to None.
-
- Returns:
- Iterator[GenerationChunk]: Generator that yields the streamed responses.
-
- Yields:
- Iterator[GenerationChunk]: Responses from the model.
- """
- return self._prepare_input_and_invoke_stream(
- prompt=prompt, stop=stop, run_manager=run_manager, **kwargs
- )
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Bedrock service model.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = llm.invoke("Tell me a joke.")
- """
-
- if self.streaming:
- completion = ""
- for chunk in self._stream(
- prompt=prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- completion += chunk.text
- return completion
-
- text, _ = self._prepare_input_and_invoke(
- prompt=prompt, stop=stop, run_manager=run_manager, **kwargs
- )
- return text
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncGenerator[GenerationChunk, None]:
- """Call out to Bedrock service with streaming.
-
- Args:
- prompt (str): The prompt to pass into the model
- stop (Optional[List[str]], optional): Stop sequences. These will
- override any stop sequences in the `model_kwargs` attribute.
- Defaults to None.
- run_manager (Optional[CallbackManagerForLLMRun], optional): Callback
- run managers used to process the output. Defaults to None.
-
- Yields:
- AsyncGenerator[GenerationChunk, None]: Generator that asynchronously yields
- the streamed responses.
- """
- async for chunk in self._aprepare_input_and_invoke_stream(
- prompt=prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- yield chunk
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Bedrock service model.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = await llm._acall("Tell me a joke.")
- """
-
- if not self.streaming:
- raise ValueError("Streaming must be set to True for async operations. ")
-
- chunks = [
- chunk.text
- async for chunk in self._astream(
- prompt=prompt, stop=stop, run_manager=run_manager, **kwargs
- )
- ]
- return "".join(chunks)
-
- def get_num_tokens(self, text: str) -> int:
- if self._model_is_anthropic:
- return get_num_tokens_anthropic(text)
- else:
- return super().get_num_tokens(text)
-
- def get_token_ids(self, text: str) -> List[int]:
- if self._model_is_anthropic:
- return get_token_ids_anthropic(text)
- else:
- return super().get_token_ids(text)
diff --git a/libs/community/langchain_community/llms/bigdl_llm.py b/libs/community/langchain_community/llms/bigdl_llm.py
deleted file mode 100644
index 59fc3e6d38..0000000000
--- a/libs/community/langchain_community/llms/bigdl_llm.py
+++ /dev/null
@@ -1,172 +0,0 @@
-import logging
-from typing import Any, Optional
-
-from langchain_core.language_models.llms import LLM
-
-from langchain_community.llms.ipex_llm import IpexLLM
-
-logger = logging.getLogger(__name__)
-
-
-class BigdlLLM(IpexLLM):
- """Wrapper around the BigdlLLM model
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import BigdlLLM
- llm = BigdlLLM.from_model_id(model_id="THUDM/chatglm-6b")
- """
-
- @classmethod
- def from_model_id(
- cls,
- model_id: str,
- model_kwargs: Optional[dict] = None,
- *,
- tokenizer_id: Optional[str] = None,
- load_in_4bit: bool = True,
- load_in_low_bit: Optional[str] = None,
- **kwargs: Any,
- ) -> LLM:
- """
- Construct object from model_id
-
- Args:
- model_id: Path for the huggingface repo id to be downloaded or
- the huggingface checkpoint folder.
- tokenizer_id: Path for the huggingface repo id to be downloaded or
- the huggingface checkpoint folder which contains the tokenizer.
- model_kwargs: Keyword arguments to pass to the model and tokenizer.
- kwargs: Extra arguments to pass to the model and tokenizer.
-
- Returns:
- An object of BigdlLLM.
- """
- logger.warning("BigdlLLM was deprecated. Please use IpexLLM instead.")
-
- try:
- from bigdl.llm.transformers import (
- AutoModel,
- AutoModelForCausalLM,
- )
- from transformers import AutoTokenizer, LlamaTokenizer
-
- except ImportError:
- raise ImportError(
- "Could not import bigdl-llm or transformers. "
- "Please install it with `pip install --pre --upgrade bigdl-llm[all]`."
- )
-
- if load_in_low_bit is not None:
- logger.warning(
- """`load_in_low_bit` option is not supported in BigdlLLM and
- is ignored. For more data types support with `load_in_low_bit`,
- use IpexLLM instead."""
- )
-
- if not load_in_4bit:
- raise ValueError(
- "BigdlLLM only supports loading in 4-bit mode, "
- "i.e. load_in_4bit = True. "
- "Please install it with `pip install --pre --upgrade bigdl-llm[all]`."
- )
-
- _model_kwargs = model_kwargs or {}
- _tokenizer_id = tokenizer_id or model_id
-
- try:
- tokenizer = AutoTokenizer.from_pretrained(_tokenizer_id, **_model_kwargs)
- except Exception:
- tokenizer = LlamaTokenizer.from_pretrained(_tokenizer_id, **_model_kwargs)
-
- try:
- model = AutoModelForCausalLM.from_pretrained(
- model_id, load_in_4bit=True, **_model_kwargs
- )
- except Exception:
- model = AutoModel.from_pretrained(
- model_id, load_in_4bit=True, **_model_kwargs
- )
-
- if "trust_remote_code" in _model_kwargs:
- _model_kwargs = {
- k: v for k, v in _model_kwargs.items() if k != "trust_remote_code"
- }
-
- return cls(
- model_id=model_id,
- model=model,
- tokenizer=tokenizer,
- model_kwargs=_model_kwargs,
- **kwargs,
- )
-
- @classmethod
- def from_model_id_low_bit(
- cls,
- model_id: str,
- model_kwargs: Optional[dict] = None,
- *,
- tokenizer_id: Optional[str] = None,
- **kwargs: Any,
- ) -> LLM:
- """
- Construct low_bit object from model_id
-
- Args:
-
- model_id: Path for the bigdl-llm transformers low-bit model folder.
- tokenizer_id: Path for the huggingface repo id or local model folder
- which contains the tokenizer.
- model_kwargs: Keyword arguments to pass to the model and tokenizer.
- kwargs: Extra arguments to pass to the model and tokenizer.
-
- Returns:
- An object of BigdlLLM.
- """
-
- logger.warning("BigdlLLM was deprecated. Please use IpexLLM instead.")
-
- try:
- from bigdl.llm.transformers import (
- AutoModel,
- AutoModelForCausalLM,
- )
- from transformers import AutoTokenizer, LlamaTokenizer
-
- except ImportError:
- raise ImportError(
- "Could not import bigdl-llm or transformers. "
- "Please install it with `pip install --pre --upgrade bigdl-llm[all]`."
- )
-
- _model_kwargs = model_kwargs or {}
- _tokenizer_id = tokenizer_id or model_id
-
- try:
- tokenizer = AutoTokenizer.from_pretrained(_tokenizer_id, **_model_kwargs)
- except Exception:
- tokenizer = LlamaTokenizer.from_pretrained(_tokenizer_id, **_model_kwargs)
-
- try:
- model = AutoModelForCausalLM.load_low_bit(model_id, **_model_kwargs)
- except Exception:
- model = AutoModel.load_low_bit(model_id, **_model_kwargs)
-
- if "trust_remote_code" in _model_kwargs:
- _model_kwargs = {
- k: v for k, v in _model_kwargs.items() if k != "trust_remote_code"
- }
-
- return cls(
- model_id=model_id,
- model=model,
- tokenizer=tokenizer,
- model_kwargs=_model_kwargs,
- **kwargs,
- )
-
- @property
- def _llm_type(self) -> str:
- return "bigdl-llm"
diff --git a/libs/community/langchain_community/llms/bittensor.py b/libs/community/langchain_community/llms/bittensor.py
deleted file mode 100644
index 3d28533f51..0000000000
--- a/libs/community/langchain_community/llms/bittensor.py
+++ /dev/null
@@ -1,174 +0,0 @@
-import http.client
-import json
-import ssl
-from typing import Any, List, Mapping, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-
-
-class NIBittensorLLM(LLM):
- """NIBittensor LLMs
-
- NIBittensorLLM is created by Neural Internet (https://neuralinternet.ai/),
- powered by Bittensor, a decentralized network full of different AI models.
-
- To analyze API_KEYS and logs of your usage visit
- https://api.neuralinternet.ai/api-keys
- https://api.neuralinternet.ai/logs
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import NIBittensorLLM
- llm = NIBittensorLLM()
- """
-
- system_prompt: Optional[str]
- """Provide system prompt that you want to supply it to model before every prompt"""
-
- top_responses: Optional[int] = 0
- """Provide top_responses to get Top N miner responses on one request.May get delayed
- Don't use in Production"""
-
- @property
- def _llm_type(self) -> str:
- return "NIBittensorLLM"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """
- Wrapper around the bittensor top miner models. Its built by Neural Internet.
-
- Call the Neural Internet's BTVEP Server and return the output.
-
- Parameters (optional):
- system_prompt(str): A system prompt defining how your model should respond.
- top_responses(int): Total top miner responses to retrieve from Bittensor
- protocol.
-
- Return:
- The generated response(s).
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import NIBittensorLLM
- llm = NIBittensorLLM(system_prompt="Act like you are programmer with \
- 5+ years of experience.")
- """
-
- # Creating HTTPS connection with SSL
- context = ssl.create_default_context()
- context.check_hostname = True
- conn = http.client.HTTPSConnection("test.neuralinternet.ai", context=context)
-
- # Sanitizing User Input before passing to API.
- if isinstance(self.top_responses, int):
- top_n = min(100, self.top_responses)
- else:
- top_n = 0
-
- default_prompt = "You are an assistant which is created by Neural Internet(NI) \
- in decentralized network named as a Bittensor."
- if self.system_prompt is None:
- system_prompt = (
- default_prompt
- + " Your task is to provide accurate response based on user prompt"
- )
- else:
- system_prompt = default_prompt + str(self.system_prompt)
-
- # Retrieving API KEY to pass into header of each request
- conn.request("GET", "/admin/api-keys/")
- api_key_response = conn.getresponse()
- api_keys_data = (
- api_key_response.read().decode("utf-8").replace("\n", "").replace("\t", "")
- )
- api_keys_json = json.loads(api_keys_data)
- api_key = api_keys_json[0]["api_key"]
-
- # Creating Header and getting top benchmark miner uids
- headers = {
- "Content-Type": "application/json",
- "Authorization": f"Bearer {api_key}",
- "Endpoint-Version": "2023-05-19",
- }
- conn.request("GET", "/top_miner_uids", headers=headers)
- miner_response = conn.getresponse()
- miner_data = (
- miner_response.read().decode("utf-8").replace("\n", "").replace("\t", "")
- )
- uids = json.loads(miner_data)
-
- # Condition for benchmark miner response
- if isinstance(uids, list) and uids and not top_n:
- for uid in uids:
- try:
- payload = json.dumps(
- {
- "uids": [uid],
- "messages": [
- {"role": "system", "content": system_prompt},
- {"role": "user", "content": prompt},
- ],
- }
- )
-
- conn.request("POST", "/chat", payload, headers)
- init_response = conn.getresponse()
- init_data = (
- init_response.read()
- .decode("utf-8")
- .replace("\n", "")
- .replace("\t", "")
- )
- init_json = json.loads(init_data)
- if "choices" not in init_json:
- continue
- reply = init_json["choices"][0]["message"]["content"]
- conn.close()
- return reply
- except Exception:
- continue
-
- # For top miner based on bittensor response
- try:
- payload = json.dumps(
- {
- "top_n": top_n,
- "messages": [
- {"role": "system", "content": system_prompt},
- {"role": "user", "content": prompt},
- ],
- }
- )
-
- conn.request("POST", "/chat", payload, headers)
- response = conn.getresponse()
- utf_string = (
- response.read().decode("utf-8").replace("\n", "").replace("\t", "")
- )
- if top_n:
- conn.close()
- return utf_string
- json_resp = json.loads(utf_string)
- reply = json_resp["choices"][0]["message"]["content"]
- conn.close()
- return reply
- except Exception as e:
- conn.request("GET", f"/error_msg?e={e}&p={prompt}", headers=headers)
- return "Sorry I am unable to provide response now, Please try again later."
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- "system_prompt": self.system_prompt,
- "top_responses": self.top_responses,
- }
diff --git a/libs/community/langchain_community/llms/cerebriumai.py b/libs/community/langchain_community/llms/cerebriumai.py
deleted file mode 100644
index b267033721..0000000000
--- a/libs/community/langchain_community/llms/cerebriumai.py
+++ /dev/null
@@ -1,113 +0,0 @@
-import logging
-from typing import Any, Dict, List, Mapping, Optional, cast
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import ConfigDict, Field, SecretStr, model_validator
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-class CerebriumAI(LLM):
- """CerebriumAI large language models.
-
- To use, you should have the ``cerebrium`` python package installed.
- You should also have the environment variable ``CEREBRIUMAI_API_KEY``
- set with your API key or pass it as a named argument in the constructor.
-
- Any parameters that are valid to be passed to the call can be passed
- in, even if not explicitly saved on this class.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import CerebriumAI
- cerebrium = CerebriumAI(endpoint_url="", cerebriumai_api_key="my-api-key")
-
- """
-
- endpoint_url: str = ""
- """model endpoint to use"""
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `create` call not
- explicitly specified."""
-
- cerebriumai_api_key: Optional[SecretStr] = None
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = set(list(cls.model_fields.keys()))
-
- extra = values.get("model_kwargs", {})
- for field_name in list(values):
- if field_name not in all_required_field_names:
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- logger.warning(
- f"""{field_name} was transferred to model_kwargs.
- Please confirm that {field_name} is what you intended."""
- )
- extra[field_name] = values.pop(field_name)
- values["model_kwargs"] = extra
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- cerebriumai_api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "cerebriumai_api_key", "CEREBRIUMAI_API_KEY")
- )
- values["cerebriumai_api_key"] = cerebriumai_api_key
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- **{"endpoint_url": self.endpoint_url},
- **{"model_kwargs": self.model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "cerebriumai"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- headers: Dict = {
- "Authorization": cast(
- SecretStr, self.cerebriumai_api_key
- ).get_secret_value(),
- "Content-Type": "application/json",
- }
- params = self.model_kwargs or {}
- payload = {"prompt": prompt, **params, **kwargs}
- response = requests.post(self.endpoint_url, json=payload, headers=headers)
- if response.status_code == 200:
- data = response.json()
- text = data["result"]
- if stop is not None:
- # I believe this is required since the stop tokens
- # are not enforced by the model parameters
- text = enforce_stop_tokens(text, stop)
- return text
- else:
- response.raise_for_status()
- return ""
diff --git a/libs/community/langchain_community/llms/chatglm.py b/libs/community/langchain_community/llms/chatglm.py
deleted file mode 100644
index c98ea1c2b1..0000000000
--- a/libs/community/langchain_community/llms/chatglm.py
+++ /dev/null
@@ -1,129 +0,0 @@
-import logging
-from typing import Any, List, Mapping, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-class ChatGLM(LLM):
- """ChatGLM LLM service.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import ChatGLM
- endpoint_url = (
- "http://127.0.0.1:8000"
- )
- ChatGLM_llm = ChatGLM(
- endpoint_url=endpoint_url
- )
- """
-
- endpoint_url: str = "http://127.0.0.1:8000/"
- """Endpoint URL to use."""
- model_kwargs: Optional[dict] = None
- """Keyword arguments to pass to the model."""
- max_token: int = 20000
- """Max token allowed to pass to the model."""
- temperature: float = 0.1
- """LLM model temperature from 0 to 10."""
- history: List[List] = []
- """History of the conversation"""
- top_p: float = 0.7
- """Top P for nucleus sampling from 0 to 1"""
- with_history: bool = False
- """Whether to use history or not"""
-
- @property
- def _llm_type(self) -> str:
- return "chat_glm"
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- **{"endpoint_url": self.endpoint_url},
- **{"model_kwargs": _model_kwargs},
- }
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to a ChatGLM LLM inference endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = chatglm_llm.invoke("Who are you?")
- """
-
- _model_kwargs = self.model_kwargs or {}
-
- # HTTP headers for authorization
- headers = {"Content-Type": "application/json"}
-
- payload = {
- "prompt": prompt,
- "temperature": self.temperature,
- "history": self.history,
- "max_length": self.max_token,
- "top_p": self.top_p,
- }
- payload.update(_model_kwargs)
- payload.update(kwargs)
-
- logger.debug(f"ChatGLM payload: {payload}")
-
- # call api
- try:
- response = requests.post(self.endpoint_url, headers=headers, json=payload)
- except requests.exceptions.RequestException as e:
- raise ValueError(f"Error raised by inference endpoint: {e}")
-
- logger.debug(f"ChatGLM response: {response}")
-
- if response.status_code != 200:
- raise ValueError(f"Failed with response: {response}")
-
- try:
- parsed_response = response.json()
-
- # Check if response content does exists
- if isinstance(parsed_response, dict):
- content_keys = "response"
- if content_keys in parsed_response:
- text = parsed_response[content_keys]
- else:
- raise ValueError(f"No content in response : {parsed_response}")
- else:
- raise ValueError(f"Unexpected response type: {parsed_response}")
-
- except requests.exceptions.JSONDecodeError as e:
- raise ValueError(
- f"Error raised during decoding response from inference endpoint: {e}."
- f"\nResponse: {response.text}"
- )
-
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
- if self.with_history:
- self.history = parsed_response["history"]
- return text
diff --git a/libs/community/langchain_community/llms/chatglm3.py b/libs/community/langchain_community/llms/chatglm3.py
deleted file mode 100644
index 796a592f4a..0000000000
--- a/libs/community/langchain_community/llms/chatglm3.py
+++ /dev/null
@@ -1,151 +0,0 @@
-import json
-import logging
-from typing import Any, List, Optional, Union
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.messages import (
- AIMessage,
- BaseMessage,
- FunctionMessage,
- HumanMessage,
- SystemMessage,
-)
-from pydantic import Field
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-HEADERS = {"Content-Type": "application/json"}
-DEFAULT_TIMEOUT = 30
-
-
-def _convert_message_to_dict(message: BaseMessage) -> dict:
- if isinstance(message, HumanMessage):
- message_dict = {"role": "user", "content": message.content}
- elif isinstance(message, AIMessage):
- message_dict = {"role": "assistant", "content": message.content}
- elif isinstance(message, SystemMessage):
- message_dict = {"role": "system", "content": message.content}
- elif isinstance(message, FunctionMessage):
- message_dict = {"role": "function", "content": message.content}
- else:
- raise ValueError(f"Got unknown type {message}")
- return message_dict
-
-
-class ChatGLM3(LLM):
- """ChatGLM3 LLM service."""
-
- model_name: str = Field(default="chatglm3-6b", alias="model")
- endpoint_url: str = "http://127.0.0.1:8000/v1/chat/completions"
- """Endpoint URL to use."""
- model_kwargs: Optional[dict] = None
- """Keyword arguments to pass to the model."""
- max_tokens: int = 20000
- """Max token allowed to pass to the model."""
- temperature: float = 0.1
- """LLM model temperature from 0 to 10."""
- top_p: float = 0.7
- """Top P for nucleus sampling from 0 to 1"""
- prefix_messages: List[BaseMessage] = Field(default_factory=list)
- """Series of messages for Chat input."""
- streaming: bool = False
- """Whether to stream the results or not."""
- http_client: Union[Any, None] = None
- timeout: int = DEFAULT_TIMEOUT
-
- @property
- def _llm_type(self) -> str:
- return "chat_glm_3"
-
- @property
- def _invocation_params(self) -> dict:
- """Get the parameters used to invoke the model."""
- params = {
- "model": self.model_name,
- "temperature": self.temperature,
- "max_tokens": self.max_tokens,
- "top_p": self.top_p,
- "stream": self.streaming,
- }
- return {**params, **(self.model_kwargs or {})}
-
- @property
- def client(self) -> Any:
- import httpx
-
- return self.http_client or httpx.Client(timeout=self.timeout)
-
- def _get_payload(self, prompt: str) -> dict:
- params = self._invocation_params
- messages = self.prefix_messages + [HumanMessage(content=prompt)]
- params.update(
- {
- "messages": [_convert_message_to_dict(m) for m in messages],
- }
- )
- return params
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to a ChatGLM3 LLM inference endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = chatglm_llm.invoke("Who are you?")
- """
- import httpx
-
- payload = self._get_payload(prompt)
- logger.debug(f"ChatGLM3 payload: {payload}")
-
- try:
- response = self.client.post(
- self.endpoint_url, headers=HEADERS, json=payload
- )
- except httpx.NetworkError as e:
- raise ValueError(f"Error raised by inference endpoint: {e}")
-
- logger.debug(f"ChatGLM3 response: {response}")
-
- if response.status_code != 200:
- raise ValueError(f"Failed with response: {response}")
-
- try:
- parsed_response = response.json()
-
- if isinstance(parsed_response, dict):
- content_keys = "choices"
- if content_keys in parsed_response:
- choices = parsed_response[content_keys]
- if len(choices):
- text = choices[0]["message"]["content"]
- else:
- raise ValueError(f"No content in response : {parsed_response}")
- else:
- raise ValueError(f"Unexpected response type: {parsed_response}")
-
- except json.JSONDecodeError as e:
- raise ValueError(
- f"Error raised during decoding response from inference endpoint: {e}."
- f"\nResponse: {response.text}"
- )
-
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
-
- return text
diff --git a/libs/community/langchain_community/llms/clarifai.py b/libs/community/langchain_community/llms/clarifai.py
deleted file mode 100644
index c7d6fcde6a..0000000000
--- a/libs/community/langchain_community/llms/clarifai.py
+++ /dev/null
@@ -1,197 +0,0 @@
-import logging
-from typing import Any, Dict, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import Generation, LLMResult
-from langchain_core.utils import pre_init
-from pydantic import ConfigDict, Field
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-EXAMPLE_URL = "https://clarifai.com/openai/chat-completion/models/GPT-4"
-
-
-class Clarifai(LLM):
- """Clarifai large language models.
-
- To use, you should have an account on the Clarifai platform,
- the ``clarifai`` python package installed, and the
- environment variable ``CLARIFAI_PAT`` set with your PAT key,
- or pass it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import Clarifai
- clarifai_llm = Clarifai(user_id=USER_ID, app_id=APP_ID, model_id=MODEL_ID)
- (or)
- clarifai_llm = Clarifai(model_url=EXAMPLE_URL)
- """
-
- model_url: Optional[str] = None
- """Model url to use."""
- model_id: Optional[str] = None
- """Model id to use."""
- model_version_id: Optional[str] = None
- """Model version id to use."""
- app_id: Optional[str] = None
- """Clarifai application id to use."""
- user_id: Optional[str] = None
- """Clarifai user id to use."""
- pat: Optional[str] = Field(default=None, exclude=True) #: :meta private:
- """Clarifai personal access token to use."""
- token: Optional[str] = Field(default=None, exclude=True) #: :meta private:
- """Clarifai session token to use."""
- model: Any = Field(default=None, exclude=True) #: :meta private:
- api_base: str = "https://api.clarifai.com"
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that we have all required info to access Clarifai
- platform and python package exists in environment."""
- try:
- from clarifai.client.model import Model
- except ImportError:
- raise ImportError(
- "Could not import clarifai python package. "
- "Please install it with `pip install clarifai`."
- )
- user_id = values.get("user_id")
- app_id = values.get("app_id")
- model_id = values.get("model_id")
- model_version_id = values.get("model_version_id")
- model_url = values.get("model_url")
- api_base = values.get("api_base")
- pat = values.get("pat")
- token = values.get("token")
-
- values["model"] = Model(
- url=model_url,
- app_id=app_id,
- user_id=user_id,
- model_version=dict(id=model_version_id),
- pat=pat,
- token=token,
- model_id=model_id,
- base_url=api_base,
- )
-
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling Clarifai API."""
- return {}
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {
- **{
- "model_url": self.model_url,
- "user_id": self.user_id,
- "app_id": self.app_id,
- "model_id": self.model_id,
- }
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "clarifai"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- inference_params: Optional[Dict[str, Any]] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Clarfai's PostModelOutputs endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = clarifai_llm.invoke("Tell me a joke.")
- """
-
- try:
- (inference_params := {}) if inference_params is None else inference_params
- predict_response = self.model.predict_by_bytes(
- bytes(prompt, "utf-8"),
- input_type="text",
- inference_params=inference_params,
- )
- text = predict_response.outputs[0].data.text.raw
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
-
- except Exception as e:
- logger.error(f"Predict failed, exception: {e}")
-
- return text
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- inference_params: Optional[Dict[str, Any]] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Run the LLM on the given prompt and input."""
-
- # TODO: add caching here.
- try:
- from clarifai.client.input import Inputs
- except ImportError:
- raise ImportError(
- "Could not import clarifai python package. "
- "Please install it with `pip install clarifai`."
- )
-
- generations = []
- batch_size = 32
- input_obj = Inputs.from_auth_helper(self.model.auth_helper)
- try:
- for i in range(0, len(prompts), batch_size):
- batch = prompts[i : i + batch_size]
- input_batch = [
- input_obj.get_text_input(input_id=str(id), raw_text=inp)
- for id, inp in enumerate(batch)
- ]
- (
- inference_params := {}
- ) if inference_params is None else inference_params
- predict_response = self.model.predict(
- inputs=input_batch, inference_params=inference_params
- )
-
- for output in predict_response.outputs:
- if stop is not None:
- text = enforce_stop_tokens(output.data.text.raw, stop)
- else:
- text = output.data.text.raw
-
- generations.append([Generation(text=text)])
-
- except Exception as e:
- logger.error(f"Predict failed, exception: {e}")
-
- return LLMResult(generations=generations)
diff --git a/libs/community/langchain_community/llms/cloudflare_workersai.py b/libs/community/langchain_community/llms/cloudflare_workersai.py
deleted file mode 100644
index 0fb6ac6053..0000000000
--- a/libs/community/langchain_community/llms/cloudflare_workersai.py
+++ /dev/null
@@ -1,128 +0,0 @@
-import json
-import logging
-from typing import Any, Dict, Iterator, List, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-
-logger = logging.getLogger(__name__)
-
-
-class CloudflareWorkersAI(LLM):
- """Cloudflare Workers AI service.
-
- To use, you must provide an API token and
- account ID to access Cloudflare Workers AI, and
- pass it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms.cloudflare_workersai import CloudflareWorkersAI
-
- my_account_id = "my_account_id"
- my_api_token = "my_secret_api_token"
- llm_model = "@cf/meta/llama-2-7b-chat-int8"
-
- cf_ai = CloudflareWorkersAI(
- account_id=my_account_id,
- api_token=my_api_token,
- model=llm_model
- )
- """ # noqa: E501
-
- account_id: str
- api_token: str
- model: str = "@cf/meta/llama-2-7b-chat-int8"
- base_url: str = "https://api.cloudflare.com/client/v4/accounts"
- streaming: bool = False
- endpoint_url: str = ""
-
- def __init__(self, **kwargs: Any) -> None:
- """Initialize the Cloudflare Workers AI class."""
- super().__init__(**kwargs)
-
- self.endpoint_url = f"{self.base_url}/{self.account_id}/ai/run/{self.model}"
-
- @property
- def _llm_type(self) -> str:
- """Return type of LLM."""
- return "cloudflare"
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Default parameters"""
- return {}
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Identifying parameters"""
- return {
- "account_id": self.account_id,
- "api_token": self.api_token,
- "model": self.model,
- "base_url": self.base_url,
- }
-
- def _call_api(self, prompt: str, params: Dict[str, Any]) -> requests.Response:
- """Call Cloudflare Workers API"""
- headers = {"Authorization": f"Bearer {self.api_token}"}
- data = {"prompt": prompt, "stream": self.streaming, **params}
- response = requests.post(
- self.endpoint_url, headers=headers, json=data, stream=self.streaming
- )
- return response
-
- def _process_response(self, response: requests.Response) -> str:
- """Process API response"""
- if response.ok:
- data = response.json()
- return data["result"]["response"]
- else:
- raise ValueError(f"Request failed with status {response.status_code}")
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- """Streaming prediction"""
- original_steaming: bool = self.streaming
- self.streaming = True
- _response_prefix_count = len("data: ")
- _response_stream_end = b"data: [DONE]"
- for chunk in self._call_api(prompt, kwargs).iter_lines():
- if chunk == _response_stream_end:
- break
- if len(chunk) > _response_prefix_count:
- try:
- data = json.loads(chunk[_response_prefix_count:])
- except Exception as e:
- logger.debug(chunk)
- raise e
- if data is not None and "response" in data:
- if run_manager:
- run_manager.on_llm_new_token(data["response"])
- yield GenerationChunk(text=data["response"])
- logger.debug("stream end")
- self.streaming = original_steaming
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Regular prediction"""
- if self.streaming:
- return "".join(
- [c.text for c in self._stream(prompt, stop, run_manager, **kwargs)]
- )
- else:
- response = self._call_api(prompt, kwargs)
- return self._process_response(response)
diff --git a/libs/community/langchain_community/llms/cohere.py b/libs/community/langchain_community/llms/cohere.py
deleted file mode 100644
index dbe15e200a..0000000000
--- a/libs/community/langchain_community/llms/cohere.py
+++ /dev/null
@@ -1,265 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Callable, Dict, List, Optional
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.load.serializable import Serializable
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import ConfigDict, Field, SecretStr
-from tenacity import (
- before_sleep_log,
- retry,
- retry_if_exception_type,
- stop_after_attempt,
- wait_exponential,
-)
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-def _create_retry_decorator(max_retries: int) -> Callable[[Any], Any]:
- import cohere
-
- # support v4 and v5
- retry_conditions = (
- retry_if_exception_type(cohere.error.CohereError)
- if hasattr(cohere, "error")
- else retry_if_exception_type(Exception)
- )
-
- min_seconds = 4
- max_seconds = 10
- # Wait 2^x * 1 second between each retry starting with
- # 4 seconds, then up to 10 seconds, then 10 seconds afterwards
- return retry(
- reraise=True,
- stop=stop_after_attempt(max_retries),
- wait=wait_exponential(multiplier=1, min=min_seconds, max=max_seconds),
- retry=retry_conditions,
- before_sleep=before_sleep_log(logger, logging.WARNING),
- )
-
-
-def completion_with_retry(llm: Cohere, **kwargs: Any) -> Any:
- """Use tenacity to retry the completion call."""
- retry_decorator = _create_retry_decorator(llm.max_retries)
-
- @retry_decorator
- def _completion_with_retry(**kwargs: Any) -> Any:
- return llm.client.generate(**kwargs)
-
- return _completion_with_retry(**kwargs)
-
-
-def acompletion_with_retry(llm: Cohere, **kwargs: Any) -> Any:
- """Use tenacity to retry the completion call."""
- retry_decorator = _create_retry_decorator(llm.max_retries)
-
- @retry_decorator
- async def _completion_with_retry(**kwargs: Any) -> Any:
- return await llm.async_client.generate(**kwargs)
-
- return _completion_with_retry(**kwargs)
-
-
-@deprecated(
- since="0.0.30", removal="1.0", alternative_import="langchain_cohere.BaseCohere"
-)
-class BaseCohere(Serializable):
- """Base class for Cohere models."""
-
- client: Any = None #: :meta private:
- async_client: Any = None #: :meta private:
- model: Optional[str] = Field(default=None)
- """Model name to use."""
-
- temperature: float = 0.75
- """A non-negative float that tunes the degree of randomness in generation."""
-
- cohere_api_key: Optional[SecretStr] = None
- """Cohere API key. If not provided, will be read from the environment variable."""
-
- stop: Optional[List[str]] = None
-
- streaming: bool = Field(default=False)
- """Whether to stream the results."""
-
- user_agent: str = "langchain"
- """Identifier for the application making the request."""
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- try:
- import cohere
- except ImportError:
- raise ImportError(
- "Could not import cohere python package. "
- "Please install it with `pip install cohere`."
- )
- else:
- values["cohere_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(values, "cohere_api_key", "COHERE_API_KEY")
- )
- client_name = values["user_agent"]
- values["client"] = cohere.Client(
- api_key=values["cohere_api_key"].get_secret_value(),
- client_name=client_name,
- )
- values["async_client"] = cohere.AsyncClient(
- api_key=values["cohere_api_key"].get_secret_value(),
- client_name=client_name,
- )
- return values
-
-
-@deprecated(since="0.1.14", removal="1.0", alternative_import="langchain_cohere.Cohere")
-class Cohere(LLM, BaseCohere):
- """Cohere large language models.
-
- To use, you should have the ``cohere`` python package installed, and the
- environment variable ``COHERE_API_KEY`` set with your API key, or pass
- it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import Cohere
-
- cohere = Cohere(model="gptd-instruct-tft", cohere_api_key="my-api-key")
- """
-
- max_tokens: int = 256
- """Denotes the number of tokens to predict per generation."""
-
- k: int = 0
- """Number of most likely tokens to consider at each step."""
-
- p: int = 1
- """Total probability mass of tokens to consider at each step."""
-
- frequency_penalty: float = 0.0
- """Penalizes repeated tokens according to frequency. Between 0 and 1."""
-
- presence_penalty: float = 0.0
- """Penalizes repeated tokens. Between 0 and 1."""
-
- truncate: Optional[str] = None
- """Specify how the client handles inputs longer than the maximum token
- length: Truncate from START, END or NONE"""
-
- max_retries: int = 10
- """Maximum number of retries to make when generating."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling Cohere API."""
- return {
- "max_tokens": self.max_tokens,
- "temperature": self.temperature,
- "k": self.k,
- "p": self.p,
- "frequency_penalty": self.frequency_penalty,
- "presence_penalty": self.presence_penalty,
- "truncate": self.truncate,
- }
-
- @property
- def lc_secrets(self) -> Dict[str, str]:
- return {"cohere_api_key": "COHERE_API_KEY"}
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {**{"model": self.model}, **self._default_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "cohere"
-
- def _invocation_params(self, stop: Optional[List[str]], **kwargs: Any) -> dict:
- params = self._default_params
- if self.stop is not None and stop is not None:
- raise ValueError("`stop` found in both the input and default params.")
- elif self.stop is not None:
- params["stop_sequences"] = self.stop
- else:
- params["stop_sequences"] = stop
- return {**params, **kwargs}
-
- def _process_response(self, response: Any, stop: Optional[List[str]]) -> str:
- text = response.generations[0].text
- # If stop tokens are provided, Cohere's endpoint returns them.
- # In order to make this consistent with other endpoints, we strip them.
- if stop:
- text = enforce_stop_tokens(text, stop)
- return text
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Cohere's generate endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = cohere("Tell me a joke.")
- """
- params = self._invocation_params(stop, **kwargs)
- response = completion_with_retry(
- self, model=self.model, prompt=prompt, **params
- )
- _stop = params.get("stop_sequences")
- return self._process_response(response, _stop)
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Async call out to Cohere's generate endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = await cohere("Tell me a joke.")
- """
- params = self._invocation_params(stop, **kwargs)
- response = await acompletion_with_retry(
- self, model=self.model, prompt=prompt, **params
- )
- _stop = params.get("stop_sequences")
- return self._process_response(response, _stop)
diff --git a/libs/community/langchain_community/llms/ctransformers.py b/libs/community/langchain_community/llms/ctransformers.py
deleted file mode 100644
index 612e6041db..0000000000
--- a/libs/community/langchain_community/llms/ctransformers.py
+++ /dev/null
@@ -1,140 +0,0 @@
-from functools import partial
-from typing import Any, Dict, List, Optional, Sequence
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import pre_init
-
-
-class CTransformers(LLM):
- """C Transformers LLM models.
-
- To use, you should have the ``ctransformers`` python package installed.
- See https://github.com/marella/ctransformers
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import CTransformers
-
- llm = CTransformers(model="/path/to/ggml-gpt-2.bin", model_type="gpt2")
- """
-
- client: Any #: :meta private:
-
- model: str
- """The path to a model file or directory or the name of a Hugging Face Hub
- model repo."""
-
- model_type: Optional[str] = None
- """The model type."""
-
- model_file: Optional[str] = None
- """The name of the model file in repo or directory."""
-
- config: Optional[Dict[str, Any]] = None
- """The config parameters.
- See https://github.com/marella/ctransformers#config"""
-
- lib: Optional[str] = None
- """The path to a shared library or one of `avx2`, `avx`, `basic`."""
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {
- "model": self.model,
- "model_type": self.model_type,
- "model_file": self.model_file,
- "config": self.config,
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "ctransformers"
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that ``ctransformers`` package is installed."""
- try:
- from ctransformers import AutoModelForCausalLM
- except ImportError:
- raise ImportError(
- "Could not import `ctransformers` package. "
- "Please install it with `pip install ctransformers`"
- )
-
- config = values["config"] or {}
- values["client"] = AutoModelForCausalLM.from_pretrained(
- values["model"],
- model_type=values["model_type"],
- model_file=values["model_file"],
- lib=values["lib"],
- **config,
- )
- return values
-
- def _call(
- self,
- prompt: str,
- stop: Optional[Sequence[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Generate text from a prompt.
-
- Args:
- prompt: The prompt to generate text from.
- stop: A list of sequences to stop generation when encountered.
-
- Returns:
- The generated text.
-
- Example:
- .. code-block:: python
-
- response = llm.invoke("Tell me a joke.")
- """
- text = []
- _run_manager = run_manager or CallbackManagerForLLMRun.get_noop_manager()
- for chunk in self.client(prompt, stop=stop, stream=True):
- text.append(chunk)
- _run_manager.on_llm_new_token(chunk, verbose=self.verbose)
- return "".join(text)
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Asynchronous Call out to CTransformers generate method.
- Very helpful when streaming (like with websockets!)
-
- Args:
- prompt: The prompt to pass into the model.
- stop: A list of strings to stop generation when encountered.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
- response = llm.invoke("Once upon a time, ")
- """
- text_callback = None
- if run_manager:
- text_callback = partial(run_manager.on_llm_new_token, verbose=self.verbose)
-
- text = ""
- for token in self.client(prompt, stop=stop, stream=True):
- if text_callback:
- await text_callback(token)
- text += token
-
- return text
diff --git a/libs/community/langchain_community/llms/ctranslate2.py b/libs/community/langchain_community/llms/ctranslate2.py
deleted file mode 100644
index bc78c7e4a4..0000000000
--- a/libs/community/langchain_community/llms/ctranslate2.py
+++ /dev/null
@@ -1,129 +0,0 @@
-from typing import Any, Dict, List, Optional, Union
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import BaseLLM
-from langchain_core.outputs import Generation, LLMResult
-from langchain_core.utils import pre_init
-from pydantic import Field
-
-
-class CTranslate2(BaseLLM):
- """CTranslate2 language model."""
-
- model_path: str = ""
- """Path to the CTranslate2 model directory."""
-
- tokenizer_name: str = ""
- """Name of the original Hugging Face model needed to load the proper tokenizer."""
-
- device: str = "cpu"
- """Device to use (possible values are: cpu, cuda, auto)."""
-
- device_index: Union[int, List[int]] = 0
- """Device IDs where to place this generator on."""
-
- compute_type: Union[str, Dict[str, str]] = "default"
- """
- Model computation type or a dictionary mapping a device name to the computation type
- (possible values are: default, auto, int8, int8_float32, int8_float16,
- int8_bfloat16, int16, float16, bfloat16, float32).
- """
-
- max_length: int = 512
- """Maximum generation length."""
-
- sampling_topk: int = 1
- """Randomly sample predictions from the top K candidates."""
-
- sampling_topp: float = 1
- """Keep the most probable tokens whose cumulative probability exceeds this value."""
-
- sampling_temperature: float = 1
- """Sampling temperature to generate more random samples."""
-
- client: Any = None #: :meta private:
-
- tokenizer: Any = None #: :meta private:
-
- ctranslate2_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """
- Holds any model parameters valid for `ctranslate2.Generator` call not
- explicitly specified.
- """
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that python package exists in environment."""
-
- try:
- import ctranslate2
- except ImportError:
- raise ImportError(
- "Could not import ctranslate2 python package. "
- "Please install it with `pip install ctranslate2`."
- )
-
- try:
- import transformers
- except ImportError:
- raise ImportError(
- "Could not import transformers python package. "
- "Please install it with `pip install transformers`."
- )
-
- values["client"] = ctranslate2.Generator(
- model_path=values["model_path"],
- device=values["device"],
- device_index=values["device_index"],
- compute_type=values["compute_type"],
- **values["ctranslate2_kwargs"],
- )
-
- values["tokenizer"] = transformers.AutoTokenizer.from_pretrained(
- values["tokenizer_name"]
- )
-
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters."""
- return {
- "max_length": self.max_length,
- "sampling_topk": self.sampling_topk,
- "sampling_topp": self.sampling_topp,
- "sampling_temperature": self.sampling_temperature,
- }
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- # build sampling parameters
- params = {**self._default_params, **kwargs}
-
- # call the model
- encoded_prompts = self.tokenizer(prompts)["input_ids"]
- tokenized_prompts = [
- self.tokenizer.convert_ids_to_tokens(encoded_prompt)
- for encoded_prompt in encoded_prompts
- ]
-
- results = self.client.generate_batch(tokenized_prompts, **params)
-
- sequences = [result.sequences_ids[0] for result in results]
- decoded_sequences = [self.tokenizer.decode(seq) for seq in sequences]
-
- generations = []
- for text in decoded_sequences:
- generations.append([Generation(text=text)])
-
- return LLMResult(generations=generations)
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "ctranslate2"
diff --git a/libs/community/langchain_community/llms/databricks.py b/libs/community/langchain_community/llms/databricks.py
deleted file mode 100644
index 4d22033406..0000000000
--- a/libs/community/langchain_community/llms/databricks.py
+++ /dev/null
@@ -1,570 +0,0 @@
-import os
-import re
-import warnings
-from abc import ABC, abstractmethod
-from typing import Any, Callable, Dict, List, Mapping, Optional
-
-import requests
-from langchain_core._api import deprecated
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models import LLM
-from pydantic import (
- BaseModel,
- ConfigDict,
- Field,
- PrivateAttr,
- model_validator,
-)
-
-__all__ = ["Databricks"]
-
-
-class _DatabricksClientBase(BaseModel, ABC):
- """A base JSON API client that talks to Databricks."""
-
- api_url: str
- api_token: str
-
- def request(self, method: str, url: str, request: Any) -> Any:
- headers = {"Authorization": f"Bearer {self.api_token}"}
- response = requests.request(
- method=method, url=url, headers=headers, json=request
- )
- # TODO: error handling and automatic retries
- if not response.ok:
- raise ValueError(f"HTTP {response.status_code} error: {response.text}")
- return response.json()
-
- def _get(self, url: str) -> Any:
- return self.request("GET", url, None)
-
- def _post(self, url: str, request: Any) -> Any:
- return self.request("POST", url, request)
-
- @abstractmethod
- def post(
- self, request: Any, transform_output_fn: Optional[Callable[..., str]] = None
- ) -> Any: ...
-
- @property
- def llm(self) -> bool:
- return False
-
-
-def _transform_completions(response: Dict[str, Any]) -> str:
- return response["choices"][0]["text"]
-
-
-def _transform_llama2_chat(response: Dict[str, Any]) -> str:
- return response["candidates"][0]["text"]
-
-
-def _transform_chat(response: Dict[str, Any]) -> str:
- return response["choices"][0]["message"]["content"]
-
-
-class _DatabricksServingEndpointClient(_DatabricksClientBase):
- """An API client that talks to a Databricks serving endpoint."""
-
- host: str
- endpoint_name: str
- databricks_uri: str
- client: Any = None
- external_or_foundation: bool = False
- task: Optional[str] = None
-
- def __init__(self, **data: Any):
- super().__init__(**data)
-
- try:
- from mlflow.deployments import get_deploy_client
-
- self.client = get_deploy_client(self.databricks_uri)
- except ImportError as e:
- raise ImportError(
- "Failed to create the client. "
- "Please install mlflow with `pip install mlflow`."
- ) from e
-
- endpoint = self.client.get_endpoint(self.endpoint_name)
- self.external_or_foundation = endpoint.get("endpoint_type", "").lower() in (
- "external_model",
- "foundation_model_api",
- )
- if self.task is None:
- self.task = endpoint.get("task")
-
- @property
- def llm(self) -> bool:
- return self.task in ("llm/v1/chat", "llm/v1/completions", "llama2/chat")
-
- @model_validator(mode="before")
- @classmethod
- def set_api_url(cls, values: Dict[str, Any]) -> Any:
- if "api_url" not in values:
- host = values["host"]
- endpoint_name = values["endpoint_name"]
- api_url = f"https://{host}/serving-endpoints/{endpoint_name}/invocations"
- values["api_url"] = api_url
- return values
-
- def post(
- self, request: Any, transform_output_fn: Optional[Callable[..., str]] = None
- ) -> Any:
- if self.external_or_foundation:
- resp = self.client.predict(endpoint=self.endpoint_name, inputs=request)
- if transform_output_fn:
- return transform_output_fn(resp)
-
- if self.task == "llm/v1/chat":
- return _transform_chat(resp)
- elif self.task == "llm/v1/completions":
- return _transform_completions(resp)
-
- return resp
- else:
- # See https://docs.databricks.com/machine-learning/model-serving/score-model-serving-endpoints.html
- wrapped_request = {"dataframe_records": [request]}
- response = self.client.predict(
- endpoint=self.endpoint_name, inputs=wrapped_request
- )
- preds = response["predictions"]
- # For a single-record query, the result is not a list.
- pred = preds[0] if isinstance(preds, list) else preds
- if self.task == "llama2/chat":
- return _transform_llama2_chat(pred)
- return transform_output_fn(pred) if transform_output_fn else pred
-
-
-class _DatabricksClusterDriverProxyClient(_DatabricksClientBase):
- """An API client that talks to a Databricks cluster driver proxy app."""
-
- host: str
- cluster_id: str
- cluster_driver_port: str
-
- @model_validator(mode="before")
- @classmethod
- def set_api_url(cls, values: Dict[str, Any]) -> Any:
- if "api_url" not in values:
- host = values["host"]
- cluster_id = values["cluster_id"]
- port = values["cluster_driver_port"]
- api_url = f"https://{host}/driver-proxy-api/o/0/{cluster_id}/{port}"
- values["api_url"] = api_url
- return values
-
- def post(
- self, request: Any, transform_output_fn: Optional[Callable[..., str]] = None
- ) -> Any:
- resp = self._post(self.api_url, request)
- return transform_output_fn(resp) if transform_output_fn else resp
-
-
-def get_repl_context() -> Any:
- """Get the notebook REPL context if running inside a Databricks notebook.
- Returns None otherwise.
- """
- try:
- from dbruntime.databricks_repl_context import get_context
-
- return get_context()
- except ImportError:
- raise ImportError(
- "Cannot access dbruntime, not running inside a Databricks notebook."
- )
-
-
-def get_default_host() -> str:
- """Get the default Databricks workspace hostname.
- Raises an error if the hostname cannot be automatically determined.
- """
- host = os.getenv("DATABRICKS_HOST")
- if not host:
- try:
- host = get_repl_context().browserHostName
- if not host:
- raise ValueError("context doesn't contain browserHostName.")
- except Exception as e:
- raise ValueError(
- "host was not set and cannot be automatically inferred. Set "
- f"environment variable 'DATABRICKS_HOST'. Received error: {e}"
- )
- # TODO: support Databricks CLI profile
- host = host.lstrip("https://").lstrip("http://").rstrip("/")
- return host
-
-
-def get_default_api_token() -> str:
- """Get the default Databricks personal access token.
- Raises an error if the token cannot be automatically determined.
- """
- if api_token := os.getenv("DATABRICKS_TOKEN"):
- return api_token
- try:
- api_token = get_repl_context().apiToken
- if not api_token:
- raise ValueError("context doesn't contain apiToken.")
- except Exception as e:
- raise ValueError(
- "api_token was not set and cannot be automatically inferred. Set "
- f"environment variable 'DATABRICKS_TOKEN'. Received error: {e}"
- )
- # TODO: support Databricks CLI profile
- return api_token
-
-
-def _is_hex_string(data: str) -> bool:
- """Checks if a data is a valid hexadecimal string using a regular expression."""
- if not isinstance(data, str):
- return False
- pattern = r"^[0-9a-fA-F]+$"
- return bool(re.match(pattern, data))
-
-
-def _load_pickled_fn_from_hex_string(
- data: str, allow_dangerous_deserialization: Optional[bool]
-) -> Callable:
- """Loads a pickled function from a hexadecimal string."""
- if not allow_dangerous_deserialization:
- raise ValueError(
- "This code relies on the pickle module. "
- "You will need to set allow_dangerous_deserialization=True "
- "if you want to opt-in to allow deserialization of data using pickle."
- "Data can be compromised by a malicious actor if "
- "not handled properly to include "
- "a malicious payload that when deserialized with "
- "pickle can execute arbitrary code on your machine."
- )
-
- try:
- import cloudpickle
- except Exception as e:
- raise ValueError(f"Please install cloudpickle>=2.0.0. Error: {e}")
-
- try:
- return cloudpickle.loads(bytes.fromhex(data)) # ignore[pickle]: explicit-opt-in
- except Exception as e:
- raise ValueError(
- f"Failed to load the pickled function from a hexadecimal string. Error: {e}"
- )
-
-
-def _pickle_fn_to_hex_string(fn: Callable) -> str:
- """Pickles a function and returns the hexadecimal string."""
- try:
- import cloudpickle
- except Exception as e:
- raise ValueError(f"Please install cloudpickle>=2.0.0. Error: {e}")
-
- try:
- return cloudpickle.dumps(fn).hex()
- except Exception as e:
- raise ValueError(f"Failed to pickle the function: {e}")
-
-
-@deprecated(
- since="0.3.3",
- removal="1.0",
- alternative_import="databricks_langchain.ChatDatabricks",
-)
-class Databricks(LLM):
- """Databricks serving endpoint or a cluster driver proxy app for LLM.
-
- It supports two endpoint types:
-
- * **Serving endpoint** (recommended for both production and development).
- We assume that an LLM was deployed to a serving endpoint.
- To wrap it as an LLM you must have "Can Query" permission to the endpoint.
- Set ``endpoint_name`` accordingly and do not set ``cluster_id`` and
- ``cluster_driver_port``.
-
- If the underlying model is a model registered by MLflow, the expected model
- signature is:
-
- * inputs::
-
- [{"name": "prompt", "type": "string"},
- {"name": "stop", "type": "list[string]"}]
-
- * outputs: ``[{"type": "string"}]``
-
- If the underlying model is an external or foundation model, the response from the
- endpoint is automatically transformed to the expected format unless
- ``transform_output_fn`` is provided.
-
- * **Cluster driver proxy app** (recommended for interactive development).
- One can load an LLM on a Databricks interactive cluster and start a local HTTP
- server on the driver node to serve the model at ``/`` using HTTP POST method
- with JSON input/output.
- Please use a port number between ``[3000, 8000]`` and let the server listen to
- the driver IP address or simply ``0.0.0.0`` instead of localhost only.
- To wrap it as an LLM you must have "Can Attach To" permission to the cluster.
- Set ``cluster_id`` and ``cluster_driver_port`` and do not set ``endpoint_name``.
- The expected server schema (using JSON schema) is:
-
- * inputs::
-
- {"type": "object",
- "properties": {
- "prompt": {"type": "string"},
- "stop": {"type": "array", "items": {"type": "string"}}},
- "required": ["prompt"]}`
-
- * outputs: ``{"type": "string"}``
-
- If the endpoint model signature is different or you want to set extra params,
- you can use `transform_input_fn` and `transform_output_fn` to apply necessary
- transformations before and after the query.
- """
-
- host: str = Field(default_factory=get_default_host)
- """Databricks workspace hostname.
- If not provided, the default value is determined by
-
- * the ``DATABRICKS_HOST`` environment variable if present, or
- * the hostname of the current Databricks workspace if running inside
- a Databricks notebook attached to an interactive cluster in "single user"
- or "no isolation shared" mode.
- """
-
- api_token: str = Field(default_factory=get_default_api_token)
- """Databricks personal access token.
- If not provided, the default value is determined by
-
- * the ``DATABRICKS_TOKEN`` environment variable if present, or
- * an automatically generated temporary token if running inside a Databricks
- notebook attached to an interactive cluster in "single user" or
- "no isolation shared" mode.
- """
-
- endpoint_name: Optional[str] = None
- """Name of the model serving endpoint.
- You must specify the endpoint name to connect to a model serving endpoint.
- You must not set both ``endpoint_name`` and ``cluster_id``.
- """
-
- cluster_id: Optional[str] = None
- """ID of the cluster if connecting to a cluster driver proxy app.
- If neither ``endpoint_name`` nor ``cluster_id`` is not provided and the code runs
- inside a Databricks notebook attached to an interactive cluster in "single user"
- or "no isolation shared" mode, the current cluster ID is used as default.
- You must not set both ``endpoint_name`` and ``cluster_id``.
- """
-
- cluster_driver_port: Optional[str] = None
- """The port number used by the HTTP server running on the cluster driver node.
- The server should listen on the driver IP address or simply ``0.0.0.0`` to connect.
- We recommend the server using a port number between ``[3000, 8000]``.
- """
-
- model_kwargs: Optional[Dict[str, Any]] = None
- """
- Deprecated. Please use ``extra_params`` instead. Extra parameters to pass to
- the endpoint.
- """
-
- transform_input_fn: Optional[Callable] = None
- """A function that transforms ``{prompt, stop, **kwargs}`` into a JSON-compatible
- request object that the endpoint accepts.
- For example, you can apply a prompt template to the input prompt.
- """
-
- transform_output_fn: Optional[Callable[..., str]] = None
- """A function that transforms the output from the endpoint to the generated text.
- """
-
- databricks_uri: str = "databricks"
- """The databricks URI. Only used when using a serving endpoint."""
-
- temperature: float = 0.0
- """The sampling temperature."""
- n: int = 1
- """The number of completion choices to generate."""
- stop: Optional[List[str]] = None
- """The stop sequence."""
- max_tokens: Optional[int] = None
- """The maximum number of tokens to generate."""
- extra_params: Dict[str, Any] = Field(default_factory=dict)
- """Any extra parameters to pass to the endpoint."""
- task: Optional[str] = None
- """The task of the endpoint. Only used when using a serving endpoint.
- If not provided, the task is automatically inferred from the endpoint.
- """
-
- allow_dangerous_deserialization: bool = False
- """Whether to allow dangerous deserialization of the data which
- involves loading data using pickle.
-
- If the data has been modified by a malicious actor, it can deliver a
- malicious payload that results in execution of arbitrary code on the target
- machine.
- """
-
- _client: _DatabricksClientBase = PrivateAttr()
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @property
- def _llm_params(self) -> Dict[str, Any]:
- params: Dict[str, Any] = {
- "temperature": self.temperature,
- "n": self.n,
- }
- if self.stop:
- params["stop"] = self.stop
- if self.max_tokens is not None:
- params["max_tokens"] = self.max_tokens
- return params
-
- @model_validator(mode="before")
- @classmethod
- def set_cluster_id(cls, values: Dict[str, Any]) -> dict:
- cluster_id = values.get("cluster_id")
- endpoint_name = values.get("endpoint_name")
- if cluster_id and endpoint_name:
- raise ValueError("Cannot set both endpoint_name and cluster_id.")
- elif endpoint_name:
- values["cluster_id"] = None
- elif cluster_id:
- pass
- else:
- try:
- if context_cluster_id := get_repl_context().clusterId:
- values["cluster_id"] = context_cluster_id
- raise ValueError("Context doesn't contain clusterId.")
- except Exception as e:
- raise ValueError(
- "Neither endpoint_name nor cluster_id was set. "
- "And the cluster_id cannot be automatically determined. Received"
- f" error: {e}"
- )
-
- cluster_driver_port = values.get("cluster_driver_port")
- if cluster_driver_port and endpoint_name:
- raise ValueError("Cannot set both endpoint_name and cluster_driver_port.")
- elif endpoint_name:
- values["cluster_driver_port"] = None
- elif cluster_driver_port is None:
- raise ValueError(
- "Must set cluster_driver_port to connect to a cluster driver."
- )
- elif int(cluster_driver_port) <= 0:
- raise ValueError(f"Invalid cluster_driver_port: {cluster_driver_port}")
- else:
- pass
-
- if model_kwargs := values.get("model_kwargs"):
- assert "prompt" not in model_kwargs, (
- "model_kwargs must not contain key 'prompt'"
- )
- assert "stop" not in model_kwargs, (
- "model_kwargs must not contain key 'stop'"
- )
- return values
-
- def __init__(self, **data: Any):
- if "transform_input_fn" in data and _is_hex_string(data["transform_input_fn"]):
- data["transform_input_fn"] = _load_pickled_fn_from_hex_string(
- data=data["transform_input_fn"],
- allow_dangerous_deserialization=data.get(
- "allow_dangerous_deserialization"
- ),
- )
- if "transform_output_fn" in data and _is_hex_string(
- data["transform_output_fn"]
- ):
- data["transform_output_fn"] = _load_pickled_fn_from_hex_string(
- data=data["transform_output_fn"],
- allow_dangerous_deserialization=data.get(
- "allow_dangerous_deserialization"
- ),
- )
-
- super().__init__(**data)
- if self.model_kwargs is not None and self.extra_params is not None:
- raise ValueError("Cannot set both extra_params and extra_params.")
- elif self.model_kwargs is not None:
- warnings.warn(
- "model_kwargs is deprecated. Please use extra_params instead.",
- DeprecationWarning,
- )
- if self.endpoint_name:
- self._client = _DatabricksServingEndpointClient(
- host=self.host,
- api_token=self.api_token,
- endpoint_name=self.endpoint_name,
- databricks_uri=self.databricks_uri,
- task=self.task,
- )
- elif self.cluster_id and self.cluster_driver_port:
- self._client = _DatabricksClusterDriverProxyClient( # type: ignore[call-arg]
- host=self.host,
- api_token=self.api_token,
- cluster_id=self.cluster_id,
- cluster_driver_port=self.cluster_driver_port,
- )
- else:
- raise ValueError(
- "Must specify either endpoint_name or cluster_id/cluster_driver_port."
- )
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Return default params."""
- return {
- "host": self.host,
- # "api_token": self.api_token, # Never save the token
- "endpoint_name": self.endpoint_name,
- "cluster_id": self.cluster_id,
- "cluster_driver_port": self.cluster_driver_port,
- "databricks_uri": self.databricks_uri,
- "model_kwargs": self.model_kwargs,
- "temperature": self.temperature,
- "n": self.n,
- "stop": self.stop,
- "max_tokens": self.max_tokens,
- "extra_params": self.extra_params,
- "task": self.task,
- "transform_input_fn": None
- if self.transform_input_fn is None
- else _pickle_fn_to_hex_string(self.transform_input_fn),
- "transform_output_fn": None
- if self.transform_output_fn is None
- else _pickle_fn_to_hex_string(self.transform_output_fn),
- }
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- return self._default_params
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "databricks"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Queries the LLM endpoint with the given prompt and stop sequence."""
-
- # TODO: support callbacks
-
- request: Dict[str, Any] = {"prompt": prompt}
- if self._client.llm:
- request.update(self._llm_params)
- request.update(self.model_kwargs or self.extra_params)
- request.update(kwargs)
- if stop:
- request["stop"] = stop
-
- if self.transform_input_fn:
- request = self.transform_input_fn(**request)
-
- return self._client.post(request, transform_output_fn=self.transform_output_fn)
diff --git a/libs/community/langchain_community/llms/deepinfra.py b/libs/community/langchain_community/llms/deepinfra.py
deleted file mode 100644
index bd0e21df52..0000000000
--- a/libs/community/langchain_community/llms/deepinfra.py
+++ /dev/null
@@ -1,246 +0,0 @@
-import json
-from typing import Any, AsyncIterator, Dict, Iterator, List, Mapping, Optional
-
-import aiohttp
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import ConfigDict
-
-from langchain_community.utilities.requests import Requests
-
-DEFAULT_MODEL_ID = "meta-llama/Meta-Llama-3-70B-Instruct"
-
-
-class DeepInfra(LLM):
- """DeepInfra models.
-
- To use, you should have the environment variable ``DEEPINFRA_API_TOKEN``
- set with your API token, or pass it as a named parameter to the
- constructor.
-
- Only supports `text-generation` and `text2text-generation` for now.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import DeepInfra
- di = DeepInfra(model_id="google/flan-t5-xl",
- deepinfra_api_token="my-api-key")
- """
-
- model_id: str = DEFAULT_MODEL_ID
- model_kwargs: Optional[Dict] = None
-
- deepinfra_api_token: Optional[str] = None
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- deepinfra_api_token = get_from_dict_or_env(
- values, "deepinfra_api_token", "DEEPINFRA_API_TOKEN"
- )
- values["deepinfra_api_token"] = deepinfra_api_token
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- **{"model_id": self.model_id},
- **{"model_kwargs": self.model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "deepinfra"
-
- def _url(self) -> str:
- return f"https://api.deepinfra.com/v1/inference/{self.model_id}"
-
- def _headers(self) -> Dict:
- return {
- "Authorization": f"bearer {self.deepinfra_api_token}",
- "Content-Type": "application/json",
- }
-
- def _body(self, prompt: str, kwargs: Any) -> Dict:
- model_kwargs = self.model_kwargs or {}
- model_kwargs = {**model_kwargs, **kwargs}
-
- return {
- "input": prompt,
- **model_kwargs,
- }
-
- def _handle_status(self, code: int, text: Any) -> None:
- if code >= 500:
- raise Exception(f"DeepInfra Server: Error {text}")
- elif code == 401:
- raise Exception("DeepInfra Server: Unauthorized")
- elif code == 403:
- raise Exception("DeepInfra Server: Unauthorized")
- elif code == 404:
- raise Exception(f"DeepInfra Server: Model not found {self.model_id}")
- elif code == 429:
- raise Exception("DeepInfra Server: Rate limit exceeded")
- elif code >= 400:
- raise ValueError(f"DeepInfra received an invalid payload: {text}")
- elif code != 200:
- raise Exception(
- f"DeepInfra returned an unexpected response with status {code}: {text}"
- )
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to DeepInfra's inference API endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = di("Tell me a joke.")
- """
-
- request = Requests(headers=self._headers())
- response = request.post(url=self._url(), data=self._body(prompt, kwargs))
-
- self._handle_status(response.status_code, response.text)
- data = response.json()
-
- return data["results"][0]["generated_text"]
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- request = Requests(headers=self._headers())
- async with request.apost(
- url=self._url(), data=self._body(prompt, kwargs)
- ) as response:
- self._handle_status(response.status, response.text)
- data = await response.json()
- return data["results"][0]["generated_text"]
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- request = Requests(headers=self._headers())
- response = request.post(
- url=self._url(), data=self._body(prompt, {**kwargs, "stream": True})
- )
- response_text = response.text
- self._handle_body_errors(response_text)
- self._handle_status(response.status_code, response.text)
- for line in _parse_stream(response.iter_lines()):
- chunk = _handle_sse_line(line)
- if chunk:
- if run_manager:
- run_manager.on_llm_new_token(chunk.text)
- yield chunk
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- request = Requests(headers=self._headers())
- async with request.apost(
- url=self._url(), data=self._body(prompt, {**kwargs, "stream": True})
- ) as response:
- response_text = await response.text()
- self._handle_body_errors(response_text)
- self._handle_status(response.status, response.text)
- async for line in _parse_stream_async(response.content):
- chunk = _handle_sse_line(line)
- if chunk:
- if run_manager:
- await run_manager.on_llm_new_token(chunk.text)
- yield chunk
-
- def _handle_body_errors(self, body: str) -> None:
- """
- Example error response:
- data: {"error_type": "validation_error",
- "error_message": "ConnectionError: ..."}
- """
- if "error" in body:
- try:
- # Remove data: prefix if present
- if body.startswith("data:"):
- body = body[len("data:") :]
- error_data = json.loads(body)
- error_message = error_data.get("error_message", "Unknown error")
-
- raise Exception(f"DeepInfra Server Error: {error_message}")
- except json.JSONDecodeError:
- raise Exception(f"DeepInfra Server: {body}")
-
-
-def _parse_stream(rbody: Iterator[bytes]) -> Iterator[str]:
- for line in rbody:
- _line = _parse_stream_helper(line)
- if _line is not None:
- yield _line
-
-
-async def _parse_stream_async(rbody: aiohttp.StreamReader) -> AsyncIterator[str]:
- async for line in rbody:
- _line = _parse_stream_helper(line)
- if _line is not None:
- yield _line
-
-
-def _parse_stream_helper(line: bytes) -> Optional[str]:
- if line and line.startswith(b"data:"):
- if line.startswith(b"data: "):
- # SSE event may be valid when it contain whitespace
- line = line[len(b"data: ") :]
- else:
- line = line[len(b"data:") :]
- if line.strip() == b"[DONE]":
- # return here will cause GeneratorExit exception in urllib3
- # and it will close http connection with TCP Reset
- return None
- else:
- return line.decode("utf-8")
- return None
-
-
-def _handle_sse_line(line: str) -> Optional[GenerationChunk]:
- try:
- obj = json.loads(line)
- return GenerationChunk(
- text=obj.get("token", {}).get("text"),
- )
- except Exception:
- return None
diff --git a/libs/community/langchain_community/llms/deepsparse.py b/libs/community/langchain_community/llms/deepsparse.py
deleted file mode 100644
index 6a743cd506..0000000000
--- a/libs/community/langchain_community/llms/deepsparse.py
+++ /dev/null
@@ -1,234 +0,0 @@
-# flake8: noqa
-from langchain_core.utils import pre_init
-from typing import Any, AsyncIterator, Dict, Iterator, List, Optional, Union
-from langchain_core.utils import pre_init
-from pydantic import root_validator
-from langchain_core.utils import pre_init
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.utils import pre_init
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import pre_init
-from langchain_community.llms.utils import enforce_stop_tokens
-from langchain_core.utils import pre_init
-from langchain_core.outputs import GenerationChunk
-
-
-class DeepSparse(LLM):
- """Neural Magic DeepSparse LLM interface.
- To use, you should have the ``deepsparse`` or ``deepsparse-nightly``
- python package installed. See https://github.com/neuralmagic/deepsparse
- This interface let's you deploy optimized LLMs straight from the
- [SparseZoo](https://sparsezoo.neuralmagic.com/?useCase=text_generation)
- Example:
- .. code-block:: python
- from langchain_community.llms import DeepSparse
- llm = DeepSparse(model="zoo:nlg/text_generation/codegen_mono-350m/pytorch/huggingface/bigpython_bigquery_thepile/base_quant-none")
- """ # noqa: E501
-
- pipeline: Any #: :meta private:
-
- model: str
- """The path to a model file or directory or the name of a SparseZoo model stub."""
-
- model_configuration: Optional[Dict[str, Any]] = None
- """Keyword arguments passed to the pipeline construction.
- Common parameters are sequence_length, prompt_sequence_length"""
-
- generation_config: Union[None, str, Dict] = None
- """GenerationConfig dictionary consisting of parameters used to control
- sequences generated for each prompt. Common parameters are:
- max_length, max_new_tokens, num_return_sequences, output_scores,
- top_p, top_k, repetition_penalty."""
-
- streaming: bool = False
- """Whether to stream the results, token by token."""
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {
- "model": self.model,
- "model_config": self.model_configuration,
- "generation_config": self.generation_config,
- "streaming": self.streaming,
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "deepsparse"
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that ``deepsparse`` package is installed."""
- try:
- from deepsparse import Pipeline
- except ImportError:
- raise ImportError(
- "Could not import `deepsparse` package. "
- "Please install it with `pip install deepsparse[llm]`"
- )
-
- model_config = values["model_configuration"] or {}
-
- values["pipeline"] = Pipeline.create(
- task="text_generation",
- model_path=values["model"],
- **model_config,
- )
- return values
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Generate text from a prompt.
- Args:
- prompt: The prompt to generate text from.
- stop: A list of strings to stop generation when encountered.
- Returns:
- The generated text.
- Example:
- .. code-block:: python
- from langchain_community.llms import DeepSparse
- llm = DeepSparse(model="zoo:nlg/text_generation/codegen_mono-350m/pytorch/huggingface/bigpython_bigquery_thepile/base_quant-none")
- llm.invoke("Tell me a joke.")
- """
- if self.streaming:
- combined_output = ""
- for chunk in self._stream(
- prompt=prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- combined_output += chunk.text
- text = combined_output
- else:
- text = (
- self.pipeline(sequences=prompt, **self.generation_config)
- .generations[0]
- .text
- )
-
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
-
- return text
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Generate text from a prompt.
- Args:
- prompt: The prompt to generate text from.
- stop: A list of strings to stop generation when encountered.
- Returns:
- The generated text.
- Example:
- .. code-block:: python
- from langchain_community.llms import DeepSparse
- llm = DeepSparse(model="zoo:nlg/text_generation/codegen_mono-350m/pytorch/huggingface/bigpython_bigquery_thepile/base_quant-none")
- llm.invoke("Tell me a joke.")
- """
- if self.streaming:
- combined_output = ""
- async for chunk in self._astream(
- prompt=prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- combined_output += chunk.text
- text = combined_output
- else:
- text = (
- self.pipeline(sequences=prompt, **self.generation_config)
- .generations[0]
- .text
- )
-
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
-
- return text
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- """Yields results objects as they are generated in real time.
- It also calls the callback manager's on_llm_new_token event with
- similar parameters to the OpenAI LLM class method of the same name.
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
- Returns:
- A generator representing the stream of tokens being generated.
- Yields:
- A dictionary like object containing a string token.
- Example:
- .. code-block:: python
- from langchain_community.llms import DeepSparse
- llm = DeepSparse(
- model="zoo:nlg/text_generation/codegen_mono-350m/pytorch/huggingface/bigpython_bigquery_thepile/base_quant-none",
- streaming=True
- )
- for chunk in llm.stream("Tell me a joke",
- stop=["'","\n"]):
- print(chunk, end='', flush=True) # noqa: T201
- """
- inference = self.pipeline(
- sequences=prompt, streaming=True, **self.generation_config
- )
- for token in inference:
- chunk = GenerationChunk(text=token.generations[0].text)
-
- if run_manager:
- run_manager.on_llm_new_token(token=chunk.text)
- yield chunk
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- """Yields results objects as they are generated in real time.
- It also calls the callback manager's on_llm_new_token event with
- similar parameters to the OpenAI LLM class method of the same name.
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
- Returns:
- A generator representing the stream of tokens being generated.
- Yields:
- A dictionary like object containing a string token.
- Example:
- .. code-block:: python
- from langchain_community.llms import DeepSparse
- llm = DeepSparse(
- model="zoo:nlg/text_generation/codegen_mono-350m/pytorch/huggingface/bigpython_bigquery_thepile/base_quant-none",
- streaming=True
- )
- for chunk in llm.stream("Tell me a joke",
- stop=["'","\n"]):
- print(chunk, end='', flush=True) # noqa: T201
- """
- inference = self.pipeline(
- sequences=prompt, streaming=True, **self.generation_config
- )
- for token in inference:
- chunk = GenerationChunk(text=token.generations[0].text)
-
- if run_manager:
- await run_manager.on_llm_new_token(token=chunk.text)
- yield chunk
diff --git a/libs/community/langchain_community/llms/edenai.py b/libs/community/langchain_community/llms/edenai.py
deleted file mode 100644
index a43e8a71be..0000000000
--- a/libs/community/langchain_community/llms/edenai.py
+++ /dev/null
@@ -1,267 +0,0 @@
-"""Wrapper around EdenAI's Generation API."""
-
-import logging
-from typing import Any, Dict, List, Literal, Optional
-
-from aiohttp import ClientSession
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from langchain_core.utils.pydantic import get_fields
-from pydantic import ConfigDict, Field, model_validator
-
-from langchain_community.llms.utils import enforce_stop_tokens
-from langchain_community.utilities.requests import Requests
-
-logger = logging.getLogger(__name__)
-
-
-class EdenAI(LLM):
- """EdenAI models.
-
- To use, you should have
- the environment variable ``EDENAI_API_KEY`` set with your API token.
- You can find your token here: https://app.edenai.run/admin/account/settings
-
- `feature` and `subfeature` are required, but any other model parameters can also be
- passed in with the format params={model_param: value, ...}
-
- for api reference check edenai documentation: http://docs.edenai.co.
- """
-
- base_url: str = "https://api.edenai.run/v2"
-
- edenai_api_key: Optional[str] = None
-
- feature: Literal["text", "image"] = "text"
- """Which generative feature to use, use text by default"""
-
- subfeature: Literal["generation"] = "generation"
- """Subfeature of above feature, use generation by default"""
-
- provider: str
- """Generative provider to use (eg: openai,stabilityai,cohere,google etc.)"""
-
- model: Optional[str] = None
- """
- model name for above provider (eg: 'gpt-3.5-turbo-instruct' for openai)
- available models are shown on https://docs.edenai.co/ under 'available providers'
- """
-
- # Optional parameters to add depending of chosen feature
- # see api reference for more infos
- temperature: Optional[float] = Field(default=None, ge=0, le=1) # for text
- max_tokens: Optional[int] = Field(default=None, ge=0) # for text
- resolution: Optional[Literal["256x256", "512x512", "1024x1024"]] = None # for image
-
- params: Dict[str, Any] = Field(default_factory=dict)
- """
- DEPRECATED: use temperature, max_tokens, resolution directly
- optional parameters to pass to api
- """
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """extra parameters"""
-
- stop_sequences: Optional[List[str]] = None
- """Stop sequences to use."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key exists in environment."""
- values["edenai_api_key"] = get_from_dict_or_env(
- values, "edenai_api_key", "EDENAI_API_KEY"
- )
- return values
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = {field.alias for field in get_fields(cls).values()}
-
- extra = values.get("model_kwargs", {})
- for field_name in list(values):
- if field_name not in all_required_field_names:
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- logger.warning(
- f"""{field_name} was transferred to model_kwargs.
- Please confirm that {field_name} is what you intended."""
- )
- extra[field_name] = values.pop(field_name)
- values["model_kwargs"] = extra
- return values
-
- @property
- def _llm_type(self) -> str:
- """Return type of model."""
- return "edenai"
-
- def _format_output(self, output: dict) -> str:
- if self.feature == "text":
- return output[self.provider]["generated_text"]
- else:
- return output[self.provider]["items"][0]["image"]
-
- @staticmethod
- def get_user_agent() -> str:
- from langchain_community import __version__
-
- return f"langchain/{__version__}"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to EdenAI's text generation endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
-
- Returns:
- json formatted str response.
- """
- stops = None
- if self.stop_sequences is not None and stop is not None:
- raise ValueError(
- "stop sequences found in both the input and default params."
- )
- elif self.stop_sequences is not None:
- stops = self.stop_sequences
- else:
- stops = stop
-
- url = f"{self.base_url}/{self.feature}/{self.subfeature}"
- headers = {
- "Authorization": f"Bearer {self.edenai_api_key}",
- "User-Agent": self.get_user_agent(),
- }
- payload: Dict[str, Any] = {
- "providers": self.provider,
- "text": prompt,
- "max_tokens": self.max_tokens,
- "temperature": self.temperature,
- "resolution": self.resolution,
- **self.params,
- **kwargs,
- "num_images": 1, # always limit to 1 (ignored for text)
- }
-
- # filter None values to not pass them to the http payload
- payload = {k: v for k, v in payload.items() if v is not None}
-
- if self.model is not None:
- payload["settings"] = {self.provider: self.model}
-
- request = Requests(headers=headers)
- response = request.post(url=url, data=payload)
-
- if response.status_code >= 500:
- raise Exception(f"EdenAI Server: Error {response.status_code}")
- elif response.status_code >= 400:
- raise ValueError(f"EdenAI received an invalid payload: {response.text}")
- elif response.status_code != 200:
- raise Exception(
- f"EdenAI returned an unexpected response with status "
- f"{response.status_code}: {response.text}"
- )
-
- data = response.json()
- provider_response = data[self.provider]
- if provider_response.get("status") == "fail":
- err_msg = provider_response.get("error", {}).get("message")
- raise Exception(err_msg)
-
- output = self._format_output(data)
-
- if stops is not None:
- output = enforce_stop_tokens(output, stops)
-
- return output
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call EdenAi model to get predictions based on the prompt.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: A list of stop words (optional).
- run_manager: A callback manager for async interaction with LLMs.
-
- Returns:
- The string generated by the model.
- """
-
- stops = None
- if self.stop_sequences is not None and stop is not None:
- raise ValueError(
- "stop sequences found in both the input and default params."
- )
- elif self.stop_sequences is not None:
- stops = self.stop_sequences
- else:
- stops = stop
-
- url = f"{self.base_url}/{self.feature}/{self.subfeature}"
- headers = {
- "Authorization": f"Bearer {self.edenai_api_key}",
- "User-Agent": self.get_user_agent(),
- }
- payload: Dict[str, Any] = {
- "providers": self.provider,
- "text": prompt,
- "max_tokens": self.max_tokens,
- "temperature": self.temperature,
- "resolution": self.resolution,
- **self.params,
- **kwargs,
- "num_images": 1, # always limit to 1 (ignored for text)
- }
-
- # filter `None` values to not pass them to the http payload as null
- payload = {k: v for k, v in payload.items() if v is not None}
-
- if self.model is not None:
- payload["settings"] = {self.provider: self.model}
-
- async with ClientSession() as session:
- async with session.post(url, json=payload, headers=headers) as response:
- if response.status >= 500:
- raise Exception(f"EdenAI Server: Error {response.status}")
- elif response.status >= 400:
- raise ValueError(
- f"EdenAI received an invalid payload: {response.text}"
- )
- elif response.status != 200:
- raise Exception(
- f"EdenAI returned an unexpected response with status "
- f"{response.status}: {response.text}"
- )
-
- response_json = await response.json()
- provider_response = response_json[self.provider]
- if provider_response.get("status") == "fail":
- err_msg = provider_response.get("error", {}).get("message")
- raise Exception(err_msg)
-
- output = self._format_output(response_json)
- if stops is not None:
- output = enforce_stop_tokens(output, stops)
-
- return output
diff --git a/libs/community/langchain_community/llms/exllamav2.py b/libs/community/langchain_community/llms/exllamav2.py
deleted file mode 100644
index 6ac9dc5832..0000000000
--- a/libs/community/langchain_community/llms/exllamav2.py
+++ /dev/null
@@ -1,200 +0,0 @@
-from typing import Any, Callable, Dict, Iterator, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import pre_init
-from pydantic import Field
-
-
-class ExLlamaV2(LLM):
- """ExllamaV2 API.
-
- - working only with GPTQ models for now.
- - Lora models are not supported yet.
-
- To use, you should have the exllamav2 library installed, and provide the
- path to the Llama model as a named parameter to the constructor.
- Check out:
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import Exllamav2
-
- llm = Exllamav2(model_path="/path/to/llama/model")
-
- #TODO:
- - Add loras support
- - Add support for custom settings
- - Add support for custom stop sequences
- """
-
- client: Any = None
- model_path: str
- exllama_cache: Any = None
- config: Any = None
- generator: Any = None
- tokenizer: Any = None
- # If settings is None, it will be used as the default settings for the model.
- # All other parameters won't be used.
- settings: Any = None
-
- # Langchain parameters
- logfunc: Callable = print
-
- stop_sequences: List[str] = Field([])
- """Sequences that immediately will stop the generator."""
-
- max_new_tokens: int = Field(150)
- """Maximum number of tokens to generate."""
-
- streaming: bool = Field(True)
- """Whether to stream the results, token by token."""
-
- verbose: bool = Field(True)
- """Whether to print debug information."""
-
- # Generator parameters
- disallowed_tokens: Optional[List[int]] = Field(None)
- """List of tokens to disallow during generation."""
-
- @pre_init
- def validate_environment(cls, values: Dict[str, Any]) -> Dict[str, Any]:
- try:
- import torch
- except ImportError as e:
- raise ImportError(
- "Unable to import torch, please install with `pip install torch`."
- ) from e
- # check if cuda is available
- if not torch.cuda.is_available():
- raise EnvironmentError("CUDA is not available. ExllamaV2 requires CUDA.")
- try:
- from exllamav2 import (
- ExLlamaV2,
- ExLlamaV2Cache,
- ExLlamaV2Config,
- ExLlamaV2Tokenizer,
- )
- from exllamav2.generator import (
- ExLlamaV2BaseGenerator,
- ExLlamaV2StreamingGenerator,
- )
- except ImportError:
- raise ImportError(
- "Could not import exllamav2 library. "
- "Please install the exllamav2 library with (cuda 12.1 is required)"
- "example : "
- "!python -m pip install https://github.com/turboderp/exllamav2/releases/download/v0.0.12/exllamav2-0.0.12+cu121-cp311-cp311-linux_x86_64.whl"
- )
-
- # Set logging function if verbose or set to empty lambda
- verbose = values["verbose"]
- if not verbose:
- values["logfunc"] = lambda *args, **kwargs: None
- logfunc = values["logfunc"]
-
- if values["settings"]:
- settings = values["settings"]
- logfunc(settings.__dict__)
- else:
- raise NotImplementedError(
- "settings is required. Custom settings are not supported yet."
- )
-
- config = ExLlamaV2Config()
- config.model_dir = values["model_path"]
- config.prepare()
-
- model = ExLlamaV2(config)
-
- exllama_cache = ExLlamaV2Cache(model, lazy=True)
- model.load_autosplit(exllama_cache)
-
- tokenizer = ExLlamaV2Tokenizer(config)
- if values["streaming"]:
- generator = ExLlamaV2StreamingGenerator(model, exllama_cache, tokenizer)
- else:
- generator = ExLlamaV2BaseGenerator(model, exllama_cache, tokenizer)
-
- # Configure the model and generator
- values["stop_sequences"] = [x.strip().lower() for x in values["stop_sequences"]]
- setattr(settings, "stop_sequences", values["stop_sequences"])
- logfunc(f"stop_sequences {values['stop_sequences']}")
-
- disallowed = values.get("disallowed_tokens")
- if disallowed:
- settings.disallow_tokens(tokenizer, disallowed)
-
- values["client"] = model
- values["generator"] = generator
- values["config"] = config
- values["tokenizer"] = tokenizer
- values["exllama_cache"] = exllama_cache
-
- return values
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "ExLlamaV2"
-
- def get_num_tokens(self, text: str) -> int:
- """Get the number of tokens present in the text."""
- return self.generator.tokenizer.num_tokens(text)
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- generator = self.generator
-
- if self.streaming:
- combined_text_output = ""
- for chunk in self._stream(
- prompt=prompt, stop=stop, run_manager=run_manager, kwargs=kwargs
- ):
- combined_text_output += str(chunk)
- return combined_text_output
- else:
- output = generator.generate_simple(
- prompt=prompt,
- gen_settings=self.settings,
- num_tokens=self.max_new_tokens,
- )
- # subtract subtext from output
- output = output[len(prompt) :]
- return output
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- input_ids = self.tokenizer.encode(prompt)
- self.generator.warmup()
- self.generator.set_stop_conditions([])
- self.generator.begin_stream(input_ids, self.settings)
-
- generated_tokens = 0
-
- while True:
- chunk, eos, _ = self.generator.stream()
- generated_tokens += 1
-
- if run_manager:
- run_manager.on_llm_new_token(
- token=chunk,
- verbose=self.verbose,
- )
- yield chunk
- if eos or generated_tokens == self.max_new_tokens:
- break
-
- return
diff --git a/libs/community/langchain_community/llms/fake.py b/libs/community/langchain_community/llms/fake.py
deleted file mode 100644
index 929fd19eb2..0000000000
--- a/libs/community/langchain_community/llms/fake.py
+++ /dev/null
@@ -1,90 +0,0 @@
-import asyncio
-import time
-from typing import Any, AsyncIterator, Iterator, List, Mapping, Optional
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models import LanguageModelInput
-from langchain_core.language_models.llms import LLM
-from langchain_core.runnables import RunnableConfig
-
-
-class FakeListLLM(LLM):
- """Fake LLM for testing purposes."""
-
- responses: List[str]
- sleep: Optional[float] = None
- i: int = 0
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "fake-list"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Return next response"""
- response = self.responses[self.i]
- if self.i < len(self.responses) - 1:
- self.i += 1
- else:
- self.i = 0
- return response
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Return next response"""
- response = self.responses[self.i]
- if self.i < len(self.responses) - 1:
- self.i += 1
- else:
- self.i = 0
- return response
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- return {"responses": self.responses}
-
-
-class FakeStreamingListLLM(FakeListLLM):
- """Fake streaming list LLM for testing purposes."""
-
- def stream(
- self,
- input: LanguageModelInput,
- config: Optional[RunnableConfig] = None,
- *,
- stop: Optional[List[str]] = None,
- **kwargs: Any,
- ) -> Iterator[str]:
- result = self.invoke(input, config)
- for c in result:
- if self.sleep is not None:
- time.sleep(self.sleep)
- yield c
-
- async def astream(
- self,
- input: LanguageModelInput,
- config: Optional[RunnableConfig] = None,
- *,
- stop: Optional[List[str]] = None,
- **kwargs: Any,
- ) -> AsyncIterator[str]:
- result = await self.ainvoke(input, config)
- for c in result:
- if self.sleep is not None:
- await asyncio.sleep(self.sleep)
- yield c
diff --git a/libs/community/langchain_community/llms/fireworks.py b/libs/community/langchain_community/llms/fireworks.py
deleted file mode 100644
index f7af9e90b8..0000000000
--- a/libs/community/langchain_community/llms/fireworks.py
+++ /dev/null
@@ -1,387 +0,0 @@
-import asyncio
-from concurrent.futures import ThreadPoolExecutor
-from typing import Any, AsyncIterator, Callable, Dict, Iterator, List, Optional, Union
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import BaseLLM, create_base_retry_decorator
-from langchain_core.outputs import Generation, GenerationChunk, LLMResult
-from langchain_core.utils import convert_to_secret_str, pre_init
-from langchain_core.utils.env import get_from_dict_or_env
-from pydantic import Field, SecretStr
-
-
-def _stream_response_to_generation_chunk(
- stream_response: Any,
-) -> GenerationChunk:
- """Convert a stream response to a generation chunk."""
- return GenerationChunk(
- text=stream_response.choices[0].text,
- generation_info=dict(
- finish_reason=stream_response.choices[0].finish_reason,
- logprobs=stream_response.choices[0].logprobs,
- ),
- )
-
-
-@deprecated(
- since="0.0.26",
- removal="1.0",
- alternative_import="langchain_fireworks.Fireworks",
-)
-class Fireworks(BaseLLM):
- """Fireworks models."""
-
- model: str = "accounts/fireworks/models/llama-v2-7b-chat"
- model_kwargs: dict = Field(
- default_factory=lambda: {
- "temperature": 0.7,
- "max_tokens": 512,
- "top_p": 1,
- }.copy()
- )
- fireworks_api_key: Optional[SecretStr] = None
- max_retries: int = 20
- batch_size: int = 20
- use_retry: bool = True
-
- @property
- def lc_secrets(self) -> Dict[str, str]:
- return {"fireworks_api_key": "FIREWORKS_API_KEY"}
-
- @classmethod
- def is_lc_serializable(cls) -> bool:
- return True
-
- @classmethod
- def get_lc_namespace(cls) -> List[str]:
- """Get the namespace of the langchain object."""
- return ["langchain", "llms", "fireworks"]
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key in environment."""
- try:
- import fireworks.client
- except ImportError as e:
- raise ImportError(
- "Could not import fireworks-ai python package. "
- "Please install it with `pip install fireworks-ai`."
- ) from e
- fireworks_api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "fireworks_api_key", "FIREWORKS_API_KEY")
- )
- fireworks.client.api_key = fireworks_api_key.get_secret_value()
- return values
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "fireworks"
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call out to Fireworks endpoint with k unique prompts.
- Args:
- prompts: The prompts to pass into the model.
- stop: Optional list of stop words to use when generating.
- Returns:
- The full LLM output.
- """
- params = {
- "model": self.model,
- **self.model_kwargs,
- }
- sub_prompts = self.get_batch_prompts(prompts)
- choices = []
- for _prompts in sub_prompts:
- response = completion_with_retry_batching(
- self,
- self.use_retry,
- prompt=_prompts,
- run_manager=run_manager,
- stop=stop,
- **params,
- )
- choices.extend(response)
-
- return self.create_llm_result(choices, prompts)
-
- async def _agenerate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call out to Fireworks endpoint async with k unique prompts."""
- params = {
- "model": self.model,
- **self.model_kwargs,
- }
- sub_prompts = self.get_batch_prompts(prompts)
- choices = []
- for _prompts in sub_prompts:
- response = await acompletion_with_retry_batching(
- self,
- self.use_retry,
- prompt=_prompts,
- run_manager=run_manager,
- stop=stop,
- **params,
- )
- choices.extend(response)
-
- return self.create_llm_result(choices, prompts)
-
- def get_batch_prompts(
- self,
- prompts: List[str],
- ) -> List[List[str]]:
- """Get the sub prompts for llm call."""
- sub_prompts = [
- prompts[i : i + self.batch_size]
- for i in range(0, len(prompts), self.batch_size)
- ]
- return sub_prompts
-
- def create_llm_result(self, choices: Any, prompts: List[str]) -> LLMResult:
- """Create the LLMResult from the choices and prompts."""
- generations = []
- for i, _ in enumerate(prompts):
- sub_choices = choices[i : (i + 1)]
- generations.append(
- [
- Generation(
- text=choice.__dict__["choices"][0].text,
- )
- for choice in sub_choices
- ]
- )
- llm_output = {"model": self.model}
- return LLMResult(generations=generations, llm_output=llm_output)
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- params = {
- "model": self.model,
- "prompt": prompt,
- "stream": True,
- **self.model_kwargs,
- }
- for stream_resp in completion_with_retry(
- self, self.use_retry, run_manager=run_manager, stop=stop, **params
- ):
- chunk = _stream_response_to_generation_chunk(stream_resp)
- if run_manager:
- run_manager.on_llm_new_token(chunk.text, chunk=chunk)
- yield chunk
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- params = {
- "model": self.model,
- "prompt": prompt,
- "stream": True,
- **self.model_kwargs,
- }
- async for stream_resp in await acompletion_with_retry_streaming(
- self, self.use_retry, run_manager=run_manager, stop=stop, **params
- ):
- chunk = _stream_response_to_generation_chunk(stream_resp)
- if run_manager:
- await run_manager.on_llm_new_token(chunk.text, chunk=chunk)
- yield chunk
-
-
-def conditional_decorator(
- condition: bool, decorator: Callable[[Any], Any]
-) -> Callable[[Any], Any]:
- """Conditionally apply a decorator.
-
- Args:
- condition: A boolean indicating whether to apply the decorator.
- decorator: A decorator function.
-
- Returns:
- A decorator function.
- """
-
- def actual_decorator(func: Callable[[Any], Any]) -> Callable[[Any], Any]:
- if condition:
- return decorator(func)
- return func
-
- return actual_decorator
-
-
-def completion_with_retry(
- llm: Fireworks,
- use_retry: bool,
- *,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
-) -> Any:
- """Use tenacity to retry the completion call."""
- import fireworks.client
-
- retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)
-
- @conditional_decorator(use_retry, retry_decorator)
- def _completion_with_retry(**kwargs: Any) -> Any:
- return fireworks.client.Completion.create(
- **kwargs,
- )
-
- return _completion_with_retry(**kwargs)
-
-
-async def acompletion_with_retry(
- llm: Fireworks,
- use_retry: bool,
- *,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
-) -> Any:
- """Use tenacity to retry the completion call."""
- import fireworks.client
-
- retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)
-
- @conditional_decorator(use_retry, retry_decorator)
- async def _completion_with_retry(**kwargs: Any) -> Any:
- return await fireworks.client.Completion.acreate(
- **kwargs,
- )
-
- return await _completion_with_retry(**kwargs)
-
-
-def completion_with_retry_batching(
- llm: Fireworks,
- use_retry: bool,
- *,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
-) -> Any:
- """Use tenacity to retry the completion call."""
- import fireworks.client
-
- prompt = kwargs["prompt"]
- del kwargs["prompt"]
-
- retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)
-
- @conditional_decorator(use_retry, retry_decorator)
- def _completion_with_retry(prompt: str) -> Any:
- return fireworks.client.Completion.create(**kwargs, prompt=prompt)
-
- def batch_sync_run() -> List:
- with ThreadPoolExecutor() as executor:
- results = list(executor.map(_completion_with_retry, prompt))
- return results
-
- return batch_sync_run()
-
-
-async def acompletion_with_retry_batching(
- llm: Fireworks,
- use_retry: bool,
- *,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
-) -> Any:
- """Use tenacity to retry the completion call."""
- import fireworks.client
-
- prompt = kwargs["prompt"]
- del kwargs["prompt"]
-
- retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)
-
- @conditional_decorator(use_retry, retry_decorator)
- async def _completion_with_retry(prompt: str) -> Any:
- return await fireworks.client.Completion.acreate(**kwargs, prompt=prompt)
-
- def run_coroutine_in_new_loop(
- coroutine_func: Any, *args: Dict, **kwargs: Dict
- ) -> Any:
- new_loop = asyncio.new_event_loop()
- try:
- asyncio.set_event_loop(new_loop)
- return new_loop.run_until_complete(coroutine_func(*args, **kwargs))
- finally:
- new_loop.close()
-
- async def batch_sync_run() -> List:
- with ThreadPoolExecutor() as executor:
- results = list(
- executor.map(
- run_coroutine_in_new_loop,
- [_completion_with_retry] * len(prompt),
- prompt,
- )
- )
- return results
-
- return await batch_sync_run()
-
-
-async def acompletion_with_retry_streaming(
- llm: Fireworks,
- use_retry: bool,
- *,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
-) -> Any:
- """Use tenacity to retry the completion call for streaming."""
- import fireworks.client
-
- retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)
-
- @conditional_decorator(use_retry, retry_decorator)
- async def _completion_with_retry(**kwargs: Any) -> Any:
- return fireworks.client.Completion.acreate(
- **kwargs,
- )
-
- return await _completion_with_retry(**kwargs)
-
-
-def _create_retry_decorator(
- llm: Fireworks,
- *,
- run_manager: Optional[
- Union[AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun]
- ] = None,
-) -> Callable[[Any], Any]:
- """Define retry mechanism."""
- import fireworks.client
-
- errors = [
- fireworks.client.error.RateLimitError,
- fireworks.client.error.InternalServerError,
- fireworks.client.error.BadGatewayError,
- fireworks.client.error.ServiceUnavailableError,
- ]
- return create_base_retry_decorator(
- error_types=errors, max_retries=llm.max_retries, run_manager=run_manager
- )
diff --git a/libs/community/langchain_community/llms/forefrontai.py b/libs/community/langchain_community/llms/forefrontai.py
deleted file mode 100644
index 8c47304438..0000000000
--- a/libs/community/langchain_community/llms/forefrontai.py
+++ /dev/null
@@ -1,119 +0,0 @@
-from typing import Any, Dict, List, Mapping, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env
-from pydantic import ConfigDict, SecretStr, model_validator
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-
-class ForefrontAI(LLM):
- """ForefrontAI large language models.
-
- To use, you should have the environment variable ``FOREFRONTAI_API_KEY``
- set with your API key.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import ForefrontAI
- forefrontai = ForefrontAI(endpoint_url="")
- """
-
- endpoint_url: str = ""
- """Model name to use."""
-
- temperature: float = 0.7
- """What sampling temperature to use."""
-
- length: int = 256
- """The maximum number of tokens to generate in the completion."""
-
- top_p: float = 1.0
- """Total probability mass of tokens to consider at each step."""
-
- top_k: int = 40
- """The number of highest probability vocabulary tokens to
- keep for top-k-filtering."""
-
- repetition_penalty: int = 1
- """Penalizes repeated tokens according to frequency."""
-
- forefrontai_api_key: SecretStr
-
- base_url: Optional[str] = None
- """Base url to use, if None decides based on model name."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key exists in environment."""
- values["forefrontai_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(values, "forefrontai_api_key", "FOREFRONTAI_API_KEY")
- )
- return values
-
- @property
- def _default_params(self) -> Mapping[str, Any]:
- """Get the default parameters for calling ForefrontAI API."""
- return {
- "temperature": self.temperature,
- "length": self.length,
- "top_p": self.top_p,
- "top_k": self.top_k,
- "repetition_penalty": self.repetition_penalty,
- }
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {**{"endpoint_url": self.endpoint_url}, **self._default_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "forefrontai"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to ForefrontAI's complete endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = ForefrontAI("Tell me a joke.")
- """
- auth_value = f"Bearer {self.forefrontai_api_key.get_secret_value()}"
- response = requests.post(
- url=self.endpoint_url,
- headers={
- "Authorization": auth_value,
- "Content-Type": "application/json",
- },
- json={"text": prompt, **self._default_params, **kwargs},
- )
- response_json = response.json()
- text = response_json["result"][0]["completion"]
- if stop is not None:
- # I believe this is required since the stop tokens
- # are not enforced by the model parameters
- text = enforce_stop_tokens(text, stop)
- return text
diff --git a/libs/community/langchain_community/llms/friendli.py b/libs/community/langchain_community/llms/friendli.py
deleted file mode 100644
index d33c80eb39..0000000000
--- a/libs/community/langchain_community/llms/friendli.py
+++ /dev/null
@@ -1,356 +0,0 @@
-from __future__ import annotations
-
-import os
-from typing import Any, AsyncIterator, Dict, Iterator, List, Optional
-
-from langchain_core.callbacks.manager import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.load.serializable import Serializable
-from langchain_core.outputs import GenerationChunk, LLMResult
-from langchain_core.utils import pre_init
-from langchain_core.utils.env import get_from_dict_or_env
-from langchain_core.utils.utils import convert_to_secret_str
-from pydantic import Field, SecretStr
-
-
-def _stream_response_to_generation_chunk(
- stream_response: Any,
-) -> GenerationChunk:
- """Convert a stream response to a generation chunk."""
- if not stream_response.get("choices", None):
- return GenerationChunk(text="")
- return GenerationChunk(
- text=stream_response.choices[0].text,
- # generation_info=dict(
- # finish_reason=stream_response.choices[0].get("finish_reason", None),
- # logprobs=stream_response.choices[0].get("logprobs", None),
- # ),
- )
-
-
-class BaseFriendli(Serializable):
- """Base class of Friendli."""
-
- # Friendli client.
- client: Any = Field(default=None, exclude=True)
- # Friendli Async client.
- async_client: Any = Field(default=None, exclude=True)
- # Model name to use.
- model: str = "meta-llama-3.1-8b-instruct"
- # Friendli personal access token to run as.
- friendli_token: Optional[SecretStr] = None
- # Friendli team ID to run as.
- friendli_team: Optional[str] = None
- # Whether to enable streaming mode.
- streaming: bool = False
- # Number between -2.0 and 2.0. Positive values penalizes tokens that have been
- # sampled, taking into account their frequency in the preceding text. This
- # penalization diminishes the model's tendency to reproduce identical lines
- # verbatim.
- frequency_penalty: Optional[float] = None
- # Number between -2.0 and 2.0. Positive values penalizes tokens that have been
- # sampled at least once in the existing text.
- presence_penalty: Optional[float] = None
- # The maximum number of tokens to generate. The length of your input tokens plus
- # `max_tokens` should not exceed the model's maximum length (e.g., 2048 for OpenAI
- # GPT-3)
- max_tokens: Optional[int] = None
- # When one of the stop phrases appears in the generation result, the API will stop
- # generation. The phrase is included in the generated result. If you are using
- # beam search, all of the active beams should contain the stop phrase to terminate
- # generation. Before checking whether a stop phrase is included in the result, the
- # phrase is converted into tokens.
- stop: Optional[List[str]] = None
- # Sampling temperature. Smaller temperature makes the generation result closer to
- # greedy, argmax (i.e., `top_k = 1`) sampling. If it is `None`, then 1.0 is used.
- temperature: Optional[float] = None
- # Tokens comprising the top `top_p` probability mass are kept for sampling. Numbers
- # between 0.0 (exclusive) and 1.0 (inclusive) are allowed. If it is `None`, then 1.0
- # is used by default.
- top_p: Optional[float] = None
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate if personal access token is provided in environment."""
- try:
- import friendli
- except ImportError as e:
- raise ImportError(
- "Could not import friendli-client python package. "
- "Please install it with `pip install friendli-client`."
- ) from e
-
- friendli_token = convert_to_secret_str(
- get_from_dict_or_env(values, "friendli_token", "FRIENDLI_TOKEN")
- )
- values["friendli_token"] = friendli_token
- friendli_token_str = friendli_token.get_secret_value()
- friendli_team = values["friendli_team"] or os.getenv("FRIENDLI_TEAM")
- values["friendli_team"] = friendli_team
- values["client"] = values["client"] or friendli.Friendli(
- token=friendli_token_str, team_id=friendli_team
- )
- values["async_client"] = values["async_client"] or friendli.AsyncFriendli(
- token=friendli_token_str, team_id=friendli_team
- )
- return values
-
-
-class Friendli(LLM, BaseFriendli):
- """Friendli LLM.
-
- ``friendli-client`` package should be installed with `pip install friendli-client`.
- You must set ``FRIENDLI_TOKEN`` environment variable or provide the value of your
- personal access token for the ``friendli_token`` argument.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import Friendli
-
- friendli = Friendli(
- model="meta-llama-3.1-8b-instruct", friendli_token="YOUR FRIENDLI TOKEN"
- )
- """
-
- @property
- def lc_secrets(self) -> Dict[str, str]:
- return {"friendli_token": "FRIENDLI_TOKEN"}
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling Friendli completions API."""
- return {
- "frequency_penalty": self.frequency_penalty,
- "presence_penalty": self.presence_penalty,
- "max_tokens": self.max_tokens,
- "stop": self.stop,
- "temperature": self.temperature,
- "top_p": self.top_p,
- }
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {"model": self.model, **self._default_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "friendli"
-
- def _get_invocation_params(
- self, stop: Optional[List[str]] = None, **kwargs: Any
- ) -> Dict[str, Any]:
- """Get the parameters used to invoke the model."""
- params = self._default_params
- if self.stop is not None and stop is not None:
- raise ValueError("`stop` found in both the input and default params.")
- elif self.stop is not None:
- params["stop"] = self.stop
- else:
- params["stop"] = stop
- return {**params, **kwargs}
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out Friendli's completions API.
-
- Args:
- prompt (str): The text prompt to generate completion for.
- stop (Optional[List[str]], optional): When one of the stop phrases appears
- in the generation result, the API will stop generation. The stop phrases
- are excluded from the result. If beam search is enabled, all of the
- active beams should contain the stop phrase to terminate generation.
- Before checking whether a stop phrase is included in the result, the
- phrase is converted into tokens. We recommend using stop_tokens because
- it is clearer. For example, after tokenization, phrases "clear" and
- " clear" can result in different token sequences due to the prepended
- space character. Defaults to None.
-
- Returns:
- str: The generated text output.
-
- Example:
- .. code-block:: python
-
- response = frienldi("Give me a recipe for the Old Fashioned cocktail.")
- """
- params = self._get_invocation_params(stop=stop, **kwargs)
- completion = self.client.completions.create(
- model=self.model, prompt=prompt, stream=False, **params
- )
- return completion.choices[0].text
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out Friendli's completions API Asynchronously.
-
- Args:
- prompt (str): The text prompt to generate completion for.
- stop (Optional[List[str]], optional): When one of the stop phrases appears
- in the generation result, the API will stop generation. The stop phrases
- are excluded from the result. If beam search is enabled, all of the
- active beams should contain the stop phrase to terminate generation.
- Before checking whether a stop phrase is included in the result, the
- phrase is converted into tokens. We recommend using stop_tokens because
- it is clearer. For example, after tokenization, phrases "clear" and
- " clear" can result in different token sequences due to the prepended
- space character. Defaults to None.
-
- Returns:
- str: The generated text output.
-
- Example:
- .. code-block:: python
-
- response = await frienldi("Tell me a joke.")
- """
- params = self._get_invocation_params(stop=stop, **kwargs)
- completion = await self.async_client.completions.create(
- model=self.model, prompt=prompt, stream=False, **params
- )
- return completion.choices[0].text
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- params = self._get_invocation_params(stop=stop, **kwargs)
- stream = self.client.completions.create(
- model=self.model, prompt=prompt, stream=True, **params
- )
- for line in stream:
- chunk = _stream_response_to_generation_chunk(line)
- if run_manager:
- run_manager.on_llm_new_token(line.text, chunk=chunk)
- yield chunk
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- params = self._get_invocation_params(stop=stop, **kwargs)
- stream = await self.async_client.completions.create(
- model=self.model, prompt=prompt, stream=True, **params
- )
- async for line in stream:
- chunk = _stream_response_to_generation_chunk(line)
- if run_manager:
- await run_manager.on_llm_new_token(line.text, chunk=chunk)
- yield chunk
-
- def _generate(
- self,
- prompts: list[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call out Friendli's completions API with k unique prompts.
-
- Args:
- prompt (str): The text prompt to generate completion for.
- stop (Optional[List[str]], optional): When one of the stop phrases appears
- in the generation result, the API will stop generation. The stop phrases
- are excluded from the result. If beam search is enabled, all of the
- active beams should contain the stop phrase to terminate generation.
- Before checking whether a stop phrase is included in the result, the
- phrase is converted into tokens. We recommend using stop_tokens because
- it is clearer. For example, after tokenization, phrases "clear" and
- " clear" can result in different token sequences due to the prepended
- space character. Defaults to None.
-
- Returns:
- str: The generated text output.
-
- Example:
- .. code-block:: python
-
- response = frienldi.generate(["Tell me a joke."])
- """
- llm_output = {"model": self.model}
- if self.streaming:
- if len(prompts) > 1:
- raise ValueError("Cannot stream results with multiple prompts.")
-
- generation: Optional[GenerationChunk] = None
- for chunk in self._stream(prompts[0], stop, run_manager, **kwargs):
- if generation is None:
- generation = chunk
- else:
- generation += chunk
- assert generation is not None
- return LLMResult(generations=[[generation]], llm_output=llm_output)
-
- llm_result = super()._generate(prompts, stop, run_manager, **kwargs)
- llm_result.llm_output = llm_output
- return llm_result
-
- async def _agenerate(
- self,
- prompts: list[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call out Friendli's completions API asynchronously with k unique prompts.
-
- Args:
- prompt (str): The text prompt to generate completion for.
- stop (Optional[List[str]], optional): When one of the stop phrases appears
- in the generation result, the API will stop generation. The stop phrases
- are excluded from the result. If beam search is enabled, all of the
- active beams should contain the stop phrase to terminate generation.
- Before checking whether a stop phrase is included in the result, the
- phrase is converted into tokens. We recommend using stop_tokens because
- it is clearer. For example, after tokenization, phrases "clear" and
- " clear" can result in different token sequences due to the prepended
- space character. Defaults to None.
-
- Returns:
- str: The generated text output.
-
- Example:
- .. code-block:: python
-
- response = await frienldi.agenerate(
- ["Give me a recipe for the Old Fashioned cocktail."]
- )
- """
- llm_output = {"model": self.model}
- if self.streaming:
- if len(prompts) > 1:
- raise ValueError("Cannot stream results with multiple prompts.")
-
- generation = None
- async for chunk in self._astream(prompts[0], stop, run_manager, **kwargs):
- if generation is None:
- generation = chunk
- else:
- generation += chunk
- assert generation is not None
- return LLMResult(generations=[[generation]], llm_output=llm_output)
-
- llm_result = await super()._agenerate(prompts, stop, run_manager, **kwargs)
- llm_result.llm_output = llm_output
- return llm_result
diff --git a/libs/community/langchain_community/llms/gigachat.py b/libs/community/langchain_community/llms/gigachat.py
deleted file mode 100644
index 0a30d7e658..0000000000
--- a/libs/community/langchain_community/llms/gigachat.py
+++ /dev/null
@@ -1,336 +0,0 @@
-from __future__ import annotations
-
-import logging
-from functools import cached_property
-from typing import TYPE_CHECKING, Any, AsyncIterator, Dict, Iterator, List, Optional
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import BaseLLM
-from langchain_core.load.serializable import Serializable
-from langchain_core.outputs import Generation, GenerationChunk, LLMResult
-from langchain_core.utils import pre_init
-from langchain_core.utils.pydantic import get_fields
-from pydantic import ConfigDict
-
-if TYPE_CHECKING:
- import gigachat
- import gigachat.models as gm
-
-logger = logging.getLogger(__name__)
-
-
-class _BaseGigaChat(Serializable):
- base_url: Optional[str] = None
- """ Base API URL """
- auth_url: Optional[str] = None
- """ Auth URL """
- credentials: Optional[str] = None
- """ Auth Token """
- scope: Optional[str] = None
- """ Permission scope for access token """
-
- access_token: Optional[str] = None
- """ Access token for GigaChat """
-
- model: Optional[str] = None
- """Model name to use."""
- user: Optional[str] = None
- """ Username for authenticate """
- password: Optional[str] = None
- """ Password for authenticate """
-
- timeout: Optional[float] = None
- """ Timeout for request """
- verify_ssl_certs: Optional[bool] = None
- """ Check certificates for all requests """
-
- ca_bundle_file: Optional[str] = None
- cert_file: Optional[str] = None
- key_file: Optional[str] = None
- key_file_password: Optional[str] = None
- # Support for connection to GigaChat through SSL certificates
-
- profanity: bool = True
- """ DEPRECATED: Check for profanity """
- profanity_check: Optional[bool] = None
- """ Check for profanity """
- streaming: bool = False
- """ Whether to stream the results or not. """
- temperature: Optional[float] = None
- """ What sampling temperature to use. """
- max_tokens: Optional[int] = None
- """ Maximum number of tokens to generate """
- use_api_for_tokens: bool = False
- """ Use GigaChat API for tokens count """
- verbose: bool = False
- """ Verbose logging """
- top_p: Optional[float] = None
- """ top_p value to use for nucleus sampling. Must be between 0.0 and 1.0 """
- repetition_penalty: Optional[float] = None
- """ The penalty applied to repeated tokens """
- update_interval: Optional[float] = None
- """ Minimum interval in seconds that elapses between sending tokens """
-
- @property
- def _llm_type(self) -> str:
- return "giga-chat-model"
-
- @property
- def lc_secrets(self) -> Dict[str, str]:
- return {
- "credentials": "GIGACHAT_CREDENTIALS",
- "access_token": "GIGACHAT_ACCESS_TOKEN",
- "password": "GIGACHAT_PASSWORD",
- "key_file_password": "GIGACHAT_KEY_FILE_PASSWORD",
- }
-
- @property
- def lc_serializable(self) -> bool:
- return True
-
- @cached_property
- def _client(self) -> gigachat.GigaChat:
- """Returns GigaChat API client"""
- import gigachat
-
- return gigachat.GigaChat(
- base_url=self.base_url,
- auth_url=self.auth_url,
- credentials=self.credentials,
- scope=self.scope,
- access_token=self.access_token,
- model=self.model,
- profanity_check=self.profanity_check,
- user=self.user,
- password=self.password,
- timeout=self.timeout,
- verify_ssl_certs=self.verify_ssl_certs,
- ca_bundle_file=self.ca_bundle_file,
- cert_file=self.cert_file,
- key_file=self.key_file,
- key_file_password=self.key_file_password,
- verbose=self.verbose,
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate authenticate data in environment and python package is installed."""
- try:
- import gigachat # noqa: F401
- except ImportError:
- raise ImportError(
- "Could not import gigachat python package. "
- "Please install it with `pip install gigachat`."
- )
- fields = set(get_fields(cls).keys())
- diff = set(values.keys()) - fields
- if diff:
- logger.warning(f"Extra fields {diff} in GigaChat class")
- if "profanity" in fields and values.get("profanity") is False:
- logger.warning(
- "'profanity' field is deprecated. Use 'profanity_check' instead."
- )
- if values.get("profanity_check") is None:
- values["profanity_check"] = values.get("profanity")
- return values
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {
- "temperature": self.temperature,
- "model": self.model,
- "profanity": self.profanity_check,
- "streaming": self.streaming,
- "max_tokens": self.max_tokens,
- "top_p": self.top_p,
- "repetition_penalty": self.repetition_penalty,
- }
-
- def tokens_count(
- self, input_: List[str], model: Optional[str] = None
- ) -> List[gm.TokensCount]:
- """Get tokens of string list"""
- return self._client.tokens_count(input_, model)
-
- async def atokens_count(
- self, input_: List[str], model: Optional[str] = None
- ) -> List[gm.TokensCount]:
- """Get tokens of strings list (async)"""
- return await self._client.atokens_count(input_, model)
-
- def get_models(self) -> gm.Models:
- """Get available models of Gigachat"""
- return self._client.get_models()
-
- async def aget_models(self) -> gm.Models:
- """Get available models of Gigachat (async)"""
- return await self._client.aget_models()
-
- def get_model(self, model: str) -> gm.Model:
- """Get info about model"""
- return self._client.get_model(model)
-
- async def aget_model(self, model: str) -> gm.Model:
- """Get info about model (async)"""
- return await self._client.aget_model(model)
-
- def get_num_tokens(self, text: str) -> int:
- """Count approximate number of tokens"""
- if self.use_api_for_tokens:
- return self.tokens_count([text])[0].tokens
- else:
- return round(len(text) / 4.6)
-
-
-class GigaChat(_BaseGigaChat, BaseLLM):
- """`GigaChat` large language models API.
-
- To use, you should pass login and password to access GigaChat API or use token.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import GigaChat
- giga = GigaChat(credentials=..., scope=..., verify_ssl_certs=False)
- """
-
- payload_role: str = "user"
-
- def _build_payload(self, messages: List[str]) -> Dict[str, Any]:
- payload: Dict[str, Any] = {
- "messages": [{"role": self.payload_role, "content": m} for m in messages],
- }
- if self.model:
- payload["model"] = self.model
- if self.profanity_check is not None:
- payload["profanity_check"] = self.profanity_check
- if self.temperature is not None:
- payload["temperature"] = self.temperature
- if self.top_p is not None:
- payload["top_p"] = self.top_p
- if self.max_tokens is not None:
- payload["max_tokens"] = self.max_tokens
- if self.repetition_penalty is not None:
- payload["repetition_penalty"] = self.repetition_penalty
- if self.update_interval is not None:
- payload["update_interval"] = self.update_interval
-
- if self.verbose:
- logger.info("Giga request: %s", payload)
-
- return payload
-
- def _create_llm_result(self, response: Any) -> LLMResult:
- generations = []
- for res in response.choices:
- finish_reason = res.finish_reason
- gen = Generation(
- text=res.message.content,
- generation_info={"finish_reason": finish_reason},
- )
- generations.append([gen])
- if finish_reason != "stop":
- logger.warning(
- "Giga generation stopped with reason: %s",
- finish_reason,
- )
- if self.verbose:
- logger.info("Giga response: %s", res.message.content)
-
- token_usage = response.usage
- llm_output = {"token_usage": token_usage, "model_name": response.model}
- return LLMResult(generations=generations, llm_output=llm_output)
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- stream: Optional[bool] = None,
- **kwargs: Any,
- ) -> LLMResult:
- should_stream = stream if stream is not None else self.streaming
- if should_stream:
- generation: Optional[GenerationChunk] = None
- stream_iter = self._stream(
- prompts[0], stop=stop, run_manager=run_manager, **kwargs
- )
- for chunk in stream_iter:
- if generation is None:
- generation = chunk
- else:
- generation += chunk
- assert generation is not None
- return LLMResult(generations=[[generation]])
-
- payload = self._build_payload(prompts)
- response = self._client.chat(payload)
-
- return self._create_llm_result(response)
-
- async def _agenerate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- stream: Optional[bool] = None,
- **kwargs: Any,
- ) -> LLMResult:
- should_stream = stream if stream is not None else self.streaming
- if should_stream:
- generation: Optional[GenerationChunk] = None
- stream_iter = self._astream(
- prompts[0], stop=stop, run_manager=run_manager, **kwargs
- )
- async for chunk in stream_iter:
- if generation is None:
- generation = chunk
- else:
- generation += chunk
- assert generation is not None
- return LLMResult(generations=[[generation]])
-
- payload = self._build_payload(prompts)
- response = await self._client.achat(payload)
-
- return self._create_llm_result(response)
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- payload = self._build_payload([prompt])
-
- for chunk in self._client.stream(payload):
- if chunk.choices:
- content = chunk.choices[0].delta.content
- if run_manager:
- run_manager.on_llm_new_token(content)
- yield GenerationChunk(text=content)
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- payload = self._build_payload([prompt])
-
- async for chunk in self._client.astream(payload):
- if chunk.choices:
- content = chunk.choices[0].delta.content
- if run_manager:
- await run_manager.on_llm_new_token(content)
- yield GenerationChunk(text=content)
-
- model_config = ConfigDict(
- extra="allow",
- )
diff --git a/libs/community/langchain_community/llms/google_palm.py b/libs/community/langchain_community/llms/google_palm.py
deleted file mode 100644
index 1d1e62bf18..0000000000
--- a/libs/community/langchain_community/llms/google_palm.py
+++ /dev/null
@@ -1,245 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Dict, Iterator, List, Optional
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models import LanguageModelInput
-from langchain_core.outputs import Generation, GenerationChunk, LLMResult
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import BaseModel, SecretStr
-
-from langchain_community.llms import BaseLLM
-from langchain_community.utilities.vertexai import create_retry_decorator
-
-
-def completion_with_retry(
- llm: GooglePalm,
- prompt: LanguageModelInput,
- is_gemini: bool = False,
- stream: bool = False,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
-) -> Any:
- """Use tenacity to retry the completion call."""
- retry_decorator = create_retry_decorator(
- llm, max_retries=llm.max_retries, run_manager=run_manager
- )
-
- @retry_decorator
- def _completion_with_retry(
- prompt: LanguageModelInput, is_gemini: bool, stream: bool, **kwargs: Any
- ) -> Any:
- generation_config = kwargs.get("generation_config", {})
- if is_gemini:
- return llm.client.generate_content(
- contents=prompt, stream=stream, generation_config=generation_config
- )
- return llm.client.generate_text(prompt=prompt, **kwargs)
-
- return _completion_with_retry(
- prompt=prompt, is_gemini=is_gemini, stream=stream, **kwargs
- )
-
-
-def _is_gemini_model(model_name: str) -> bool:
- return "gemini" in model_name
-
-
-def _strip_erroneous_leading_spaces(text: str) -> str:
- """Strip erroneous leading spaces from text.
-
- The PaLM API will sometimes erroneously return a single leading space in all
- lines > 1. This function strips that space.
- """
- has_leading_space = all(not line or line[0] == " " for line in text.split("\n")[1:])
- if has_leading_space:
- return text.replace("\n ", "\n")
- else:
- return text
-
-
-@deprecated("0.0.12", alternative_import="langchain_google_genai.GoogleGenerativeAI")
-class GooglePalm(BaseLLM, BaseModel):
- """
- DEPRECATED: Use `langchain_google_genai.GoogleGenerativeAI` instead.
-
- Google PaLM models.
- """
-
- client: Any #: :meta private:
- google_api_key: Optional[SecretStr]
- model_name: str = "models/text-bison-001"
- """Model name to use."""
- temperature: float = 0.7
- """Run inference with this temperature. Must be in the closed interval
- [0.0, 1.0]."""
- top_p: Optional[float] = None
- """Decode using nucleus sampling: consider the smallest set of tokens whose
- probability sum is at least top_p. Must be in the closed interval [0.0, 1.0]."""
- top_k: Optional[int] = None
- """Decode using top-k sampling: consider the set of top_k most probable tokens.
- Must be positive."""
- max_output_tokens: Optional[int] = None
- """Maximum number of tokens to include in a candidate. Must be greater than zero.
- If unset, will default to 64."""
- n: int = 1
- """Number of chat completions to generate for each prompt. Note that the API may
- not return the full n completions if duplicates are generated."""
- max_retries: int = 6
- """The maximum number of retries to make when generating."""
-
- @property
- def is_gemini(self) -> bool:
- """Returns whether a model is belongs to a Gemini family or not."""
- return _is_gemini_model(self.model_name)
-
- @property
- def lc_secrets(self) -> Dict[str, str]:
- return {"google_api_key": "GOOGLE_API_KEY"}
-
- @classmethod
- def is_lc_serializable(self) -> bool:
- return True
-
- @classmethod
- def get_lc_namespace(cls) -> List[str]:
- """Get the namespace of the langchain object."""
- return ["langchain", "llms", "google_palm"]
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate api key, python package exists."""
- google_api_key = get_from_dict_or_env(
- values, "google_api_key", "GOOGLE_API_KEY"
- )
- model_name = values["model_name"]
- try:
- import google.generativeai as genai
-
- if isinstance(google_api_key, SecretStr):
- google_api_key = google_api_key.get_secret_value()
-
- genai.configure(api_key=google_api_key)
-
- if _is_gemini_model(model_name):
- values["client"] = genai.GenerativeModel(model_name=model_name)
- else:
- values["client"] = genai
- except ImportError:
- raise ImportError(
- "Could not import google-generativeai python package. "
- "Please install it with `pip install google-generativeai`."
- )
-
- if values["temperature"] is not None and not 0 <= values["temperature"] <= 1:
- raise ValueError("temperature must be in the range [0.0, 1.0]")
-
- if values["top_p"] is not None and not 0 <= values["top_p"] <= 1:
- raise ValueError("top_p must be in the range [0.0, 1.0]")
-
- if values["top_k"] is not None and values["top_k"] <= 0:
- raise ValueError("top_k must be positive")
-
- if values["max_output_tokens"] is not None and values["max_output_tokens"] <= 0:
- raise ValueError("max_output_tokens must be greater than zero")
-
- return values
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- generations: List[List[Generation]] = []
- generation_config = {
- "stop_sequences": stop,
- "temperature": self.temperature,
- "top_p": self.top_p,
- "top_k": self.top_k,
- "max_output_tokens": self.max_output_tokens,
- "candidate_count": self.n,
- }
- for prompt in prompts:
- if self.is_gemini:
- res = completion_with_retry(
- self,
- prompt=prompt,
- stream=False,
- is_gemini=True,
- run_manager=run_manager,
- generation_config=generation_config,
- )
- candidates = [
- "".join([p.text for p in c.content.parts]) for c in res.candidates
- ]
- generations.append([Generation(text=c) for c in candidates])
- else:
- res = completion_with_retry(
- self,
- model=self.model_name,
- prompt=prompt,
- stream=False,
- is_gemini=False,
- run_manager=run_manager,
- **generation_config,
- )
- prompt_generations = []
- for candidate in res.candidates:
- raw_text = candidate["output"]
- stripped_text = _strip_erroneous_leading_spaces(raw_text)
- prompt_generations.append(Generation(text=stripped_text))
- generations.append(prompt_generations)
-
- return LLMResult(generations=generations)
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- generation_config = kwargs.get("generation_config", {})
- if stop:
- generation_config["stop_sequences"] = stop
- for stream_resp in completion_with_retry(
- self,
- prompt,
- stream=True,
- is_gemini=True,
- run_manager=run_manager,
- generation_config=generation_config,
- **kwargs,
- ):
- chunk = GenerationChunk(text=stream_resp.text)
- if run_manager:
- run_manager.on_llm_new_token(
- stream_resp.text,
- chunk=chunk,
- verbose=self.verbose,
- )
- yield chunk
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "google_palm"
-
- def get_num_tokens(self, text: str) -> int:
- """Get the number of tokens present in the text.
-
- Useful for checking if an input will fit in a model's context window.
-
- Args:
- text: The string input to tokenize.
-
- Returns:
- The integer number of tokens in the text.
- """
- if self.is_gemini:
- raise ValueError("Counting tokens is not yet supported!")
- result = self.client.count_text_tokens(model=self.model_name, prompt=text)
- return result["token_count"]
diff --git a/libs/community/langchain_community/llms/gooseai.py b/libs/community/langchain_community/llms/gooseai.py
deleted file mode 100644
index 282bd7d8f8..0000000000
--- a/libs/community/langchain_community/llms/gooseai.py
+++ /dev/null
@@ -1,152 +0,0 @@
-import logging
-from typing import Any, Dict, List, Mapping, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import (
- convert_to_secret_str,
- get_from_dict_or_env,
- get_pydantic_field_names,
-)
-from pydantic import ConfigDict, Field, SecretStr, model_validator
-
-logger = logging.getLogger(__name__)
-
-
-class GooseAI(LLM):
- """GooseAI large language models.
-
- To use, you should have the ``openai`` python package installed, and the
- environment variable ``GOOSEAI_API_KEY`` set with your API key.
-
- Any parameters that are valid to be passed to the openai.create call can be passed
- in, even if not explicitly saved on this class.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import GooseAI
- gooseai = GooseAI(model_name="gpt-neo-20b")
-
- """
-
- client: Any = None
-
- model_name: str = "gpt-neo-20b"
- """Model name to use"""
-
- temperature: float = 0.7
- """What sampling temperature to use"""
-
- max_tokens: int = 256
- """The maximum number of tokens to generate in the completion.
- -1 returns as many tokens as possible given the prompt and
- the models maximal context size."""
-
- top_p: float = 1
- """Total probability mass of tokens to consider at each step."""
-
- min_tokens: int = 1
- """The minimum number of tokens to generate in the completion."""
-
- frequency_penalty: float = 0
- """Penalizes repeated tokens according to frequency."""
-
- presence_penalty: float = 0
- """Penalizes repeated tokens."""
-
- n: int = 1
- """How many completions to generate for each prompt."""
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `create` call not explicitly specified."""
-
- logit_bias: Optional[Dict[str, float]] = Field(default_factory=dict) # type: ignore[arg-type]
- """Adjust the probability of specific tokens being generated."""
-
- gooseai_api_key: Optional[SecretStr] = None
-
- model_config = ConfigDict(
- extra="ignore",
- )
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = get_pydantic_field_names(cls)
-
- extra = values.get("model_kwargs", {})
- for field_name in list(values):
- if field_name not in all_required_field_names:
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- logger.warning(
- f"""WARNING! {field_name} is not default parameter.
- {field_name} was transferred to model_kwargs.
- Please confirm that {field_name} is what you intended."""
- )
- extra[field_name] = values.pop(field_name)
- values["model_kwargs"] = extra
-
- gooseai_api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "gooseai_api_key", "GOOSEAI_API_KEY")
- )
- values["gooseai_api_key"] = gooseai_api_key
- try:
- import openai
-
- openai.api_key = gooseai_api_key.get_secret_value()
- openai.api_base = "https://api.goose.ai/v1"
- values["client"] = openai.Completion
- except ImportError:
- raise ImportError(
- "Could not import openai python package. "
- "Please install it with `pip install openai`."
- )
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling GooseAI API."""
- normal_params = {
- "temperature": self.temperature,
- "max_tokens": self.max_tokens,
- "top_p": self.top_p,
- "min_tokens": self.min_tokens,
- "frequency_penalty": self.frequency_penalty,
- "presence_penalty": self.presence_penalty,
- "n": self.n,
- "logit_bias": self.logit_bias,
- }
- return {**normal_params, **self.model_kwargs}
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {**{"model_name": self.model_name}, **self._default_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "gooseai"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call the GooseAI API."""
- params = self._default_params
- if stop is not None:
- if "stop" in params:
- raise ValueError("`stop` found in both the input and default params.")
- params["stop"] = stop
-
- params = {**params, **kwargs}
-
- response = self.client.create(engine=self.model_name, prompt=prompt, **params)
- text = response.choices[0].text
- return text
diff --git a/libs/community/langchain_community/llms/gpt4all.py b/libs/community/langchain_community/llms/gpt4all.py
deleted file mode 100644
index 85c47ac691..0000000000
--- a/libs/community/langchain_community/llms/gpt4all.py
+++ /dev/null
@@ -1,213 +0,0 @@
-from functools import partial
-from typing import Any, Dict, List, Mapping, Optional, Set
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import pre_init
-from pydantic import ConfigDict, Field
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-
-class GPT4All(LLM):
- """GPT4All language models.
-
- To use, you should have the ``gpt4all`` python package installed, the
- pre-trained model file, and the model's config information.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import GPT4All
- model = GPT4All(model="./models/gpt4all-model.bin", n_threads=8)
-
- # Simplest invocation
- response = model.invoke("Once upon a time, ")
- """
-
- model: str
- """Path to the pre-trained GPT4All model file."""
-
- backend: Optional[str] = Field(None, alias="backend")
-
- max_tokens: int = Field(200, alias="max_tokens")
- """Token context window."""
-
- n_parts: int = Field(-1, alias="n_parts")
- """Number of parts to split the model into.
- If -1, the number of parts is automatically determined."""
-
- seed: int = Field(0, alias="seed")
- """Seed. If -1, a random seed is used."""
-
- f16_kv: bool = Field(False, alias="f16_kv")
- """Use half-precision for key/value cache."""
-
- logits_all: bool = Field(False, alias="logits_all")
- """Return logits for all tokens, not just the last token."""
-
- vocab_only: bool = Field(False, alias="vocab_only")
- """Only load the vocabulary, no weights."""
-
- use_mlock: bool = Field(False, alias="use_mlock")
- """Force system to keep model in RAM."""
-
- embedding: bool = Field(False, alias="embedding")
- """Use embedding mode only."""
-
- n_threads: Optional[int] = Field(4, alias="n_threads")
- """Number of threads to use."""
-
- n_predict: Optional[int] = 256
- """The maximum number of tokens to generate."""
-
- temp: Optional[float] = 0.7
- """The temperature to use for sampling."""
-
- top_p: Optional[float] = 0.1
- """The top-p value to use for sampling."""
-
- top_k: Optional[int] = 40
- """The top-k value to use for sampling."""
-
- echo: Optional[bool] = False
- """Whether to echo the prompt."""
-
- stop: Optional[List[str]] = []
- """A list of strings to stop generation when encountered."""
-
- repeat_last_n: Optional[int] = 64
- "Last n tokens to penalize"
-
- repeat_penalty: Optional[float] = 1.18
- """The penalty to apply to repeated tokens."""
-
- n_batch: int = Field(8, alias="n_batch")
- """Batch size for prompt processing."""
-
- streaming: bool = False
- """Whether to stream the results or not."""
-
- allow_download: bool = False
- """If model does not exist in ~/.cache/gpt4all/, download it."""
-
- device: Optional[str] = Field("cpu", alias="device")
- """Device name: cpu, gpu, nvidia, intel, amd or DeviceName."""
-
- client: Any = None #: :meta private:
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @staticmethod
- def _model_param_names() -> Set[str]:
- return {
- "max_tokens",
- "n_predict",
- "top_k",
- "top_p",
- "temp",
- "n_batch",
- "repeat_penalty",
- "repeat_last_n",
- "streaming",
- }
-
- def _default_params(self) -> Dict[str, Any]:
- return {
- "max_tokens": self.max_tokens,
- "n_predict": self.n_predict,
- "top_k": self.top_k,
- "top_p": self.top_p,
- "temp": self.temp,
- "n_batch": self.n_batch,
- "repeat_penalty": self.repeat_penalty,
- "repeat_last_n": self.repeat_last_n,
- "streaming": self.streaming,
- }
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that the python package exists in the environment."""
- try:
- from gpt4all import GPT4All as GPT4AllModel
- except ImportError:
- raise ImportError(
- "Could not import gpt4all python package. "
- "Please install it with `pip install gpt4all`."
- )
-
- full_path = values["model"]
- model_path, delimiter, model_name = full_path.rpartition("/")
- model_path += delimiter
-
- values["client"] = GPT4AllModel(
- model_name,
- model_path=model_path or None,
- model_type=values["backend"],
- allow_download=values["allow_download"],
- device=values["device"],
- )
- if values["n_threads"] is not None:
- # set n_threads
- values["client"].model.set_thread_count(values["n_threads"])
-
- try:
- values["backend"] = values["client"].model_type
- except AttributeError:
- # The below is for compatibility with GPT4All Python bindings <= 0.2.3.
- values["backend"] = values["client"].model.model_type
-
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- "model": self.model,
- **self._default_params(),
- **{
- k: v for k, v in self.__dict__.items() if k in self._model_param_names()
- },
- }
-
- @property
- def _llm_type(self) -> str:
- """Return the type of llm."""
- return "gpt4all"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- r"""Call out to GPT4All's generate method.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: A list of strings to stop generation when encountered.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- prompt = "Once upon a time, "
- response = model.invoke(prompt, n_predict=55)
- """
- text_callback = None
- if run_manager:
- text_callback = partial(run_manager.on_llm_new_token, verbose=self.verbose)
- text = ""
- params = {**self._default_params(), **kwargs}
- for token in self.client.generate(prompt, **params):
- if text_callback:
- text_callback(token)
- text += token
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
- return text
diff --git a/libs/community/langchain_community/llms/gradient_ai.py b/libs/community/langchain_community/llms/gradient_ai.py
deleted file mode 100644
index ee088808e9..0000000000
--- a/libs/community/langchain_community/llms/gradient_ai.py
+++ /dev/null
@@ -1,407 +0,0 @@
-import asyncio
-import logging
-from concurrent.futures import ThreadPoolExecutor
-from typing import Any, Dict, List, Mapping, Optional, Sequence, TypedDict
-
-import aiohttp
-import requests
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import BaseLLM
-from langchain_core.outputs import Generation, LLMResult
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import ConfigDict, Field, model_validator
-from typing_extensions import Self
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-
-class TrainResult(TypedDict):
- """Train result."""
-
- loss: float
-
-
-class GradientLLM(BaseLLM):
- """Gradient.ai LLM Endpoints.
-
- GradientLLM is a class to interact with LLMs on gradient.ai
-
- To use, set the environment variable ``GRADIENT_ACCESS_TOKEN`` with your
- API token and ``GRADIENT_WORKSPACE_ID`` for your gradient workspace,
- or alternatively provide them as keywords to the constructor of this class.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import GradientLLM
- GradientLLM(
- model="99148c6d-c2a0-4fbe-a4a7-e7c05bdb8a09_base_ml_model",
- model_kwargs={
- "max_generated_token_count": 128,
- "temperature": 0.75,
- "top_p": 0.95,
- "top_k": 20,
- "stop": [],
- },
- gradient_workspace_id="12345614fc0_workspace",
- gradient_access_token="gradientai-access_token",
- )
-
- """
-
- model_id: str = Field(alias="model", min_length=2)
- "Underlying gradient.ai model id (base or fine-tuned)."
-
- gradient_workspace_id: Optional[str] = None
- "Underlying gradient.ai workspace_id."
-
- gradient_access_token: Optional[str] = None
- """gradient.ai API Token, which can be generated by going to
- https://auth.gradient.ai/select-workspace
- and selecting "Access tokens" under the profile drop-down.
- """
-
- model_kwargs: Optional[dict] = None
- """Keyword arguments to pass to the model."""
-
- gradient_api_url: str = "https://api.gradient.ai/api"
- """Endpoint URL to use."""
-
- aiosession: Optional[aiohttp.ClientSession] = None #: :meta private:
- """ClientSession, private, subject to change in upcoming releases."""
-
- # LLM call kwargs
- model_config = ConfigDict(
- populate_by_name=True,
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and python package exists in environment."""
-
- values["gradient_access_token"] = get_from_dict_or_env(
- values, "gradient_access_token", "GRADIENT_ACCESS_TOKEN"
- )
- values["gradient_workspace_id"] = get_from_dict_or_env(
- values, "gradient_workspace_id", "GRADIENT_WORKSPACE_ID"
- )
-
- values["gradient_api_url"] = get_from_dict_or_env(
- values, "gradient_api_url", "GRADIENT_API_URL"
- )
- return values
-
- @model_validator(mode="after")
- def post_init(self) -> Self:
- """Post init validation."""
- # Can be most to post_init_validation
- try:
- import gradientai # noqa
- except ImportError:
- logging.warning(
- "DeprecationWarning: `GradientLLM` will use "
- "`pip install gradientai` in future releases of langchain."
- )
- except Exception:
- pass
-
- # Can be most to post_init_validation
- if self.gradient_access_token is None or len(self.gradient_access_token) < 10:
- raise ValueError("env variable `GRADIENT_ACCESS_TOKEN` must be set")
-
- if self.gradient_workspace_id is None or len(self.gradient_access_token) < 3:
- raise ValueError("env variable `GRADIENT_WORKSPACE_ID` must be set")
-
- if self.model_kwargs:
- kw = self.model_kwargs
- if not 0 <= kw.get("temperature", 0.5) <= 1:
- raise ValueError("`temperature` must be in the range [0.0, 1.0]")
-
- if not 0 <= kw.get("top_p", 0.5) <= 1:
- raise ValueError("`top_p` must be in the range [0.0, 1.0]")
-
- if 0 >= kw.get("top_k", 0.5):
- raise ValueError("`top_k` must be positive")
-
- if 0 >= kw.get("max_generated_token_count", 1):
- raise ValueError("`max_generated_token_count` must be positive")
-
- return self
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- **{"gradient_api_url": self.gradient_api_url},
- **{"model_kwargs": _model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "gradient"
-
- def _kwargs_post_fine_tune_request(
- self, inputs: Sequence[str], kwargs: Mapping[str, Any]
- ) -> Mapping[str, Any]:
- """Build the kwargs for the Post request, used by sync
-
- Args:
- prompt (str): prompt used in query
- kwargs (dict): model kwargs in payload
-
- Returns:
- Dict[str, Union[str,dict]]: _description_
- """
- _model_kwargs = self.model_kwargs or {}
- _params = {**_model_kwargs, **kwargs}
-
- multipliers = _params.get("multipliers", None)
-
- return dict(
- url=f"{self.gradient_api_url}/models/{self.model_id}/fine-tune",
- headers={
- "authorization": f"Bearer {self.gradient_access_token}",
- "x-gradient-workspace-id": f"{self.gradient_workspace_id}",
- "accept": "application/json",
- "content-type": "application/json",
- },
- json=dict(
- samples=(
- tuple(
- {
- "inputs": input,
- }
- for input in inputs
- )
- if multipliers is None
- else tuple(
- {
- "inputs": input,
- "fineTuningParameters": {
- "multiplier": multiplier,
- },
- }
- for input, multiplier in zip(inputs, multipliers)
- )
- ),
- ),
- )
-
- def _kwargs_post_request(
- self, prompt: str, kwargs: Mapping[str, Any]
- ) -> Mapping[str, Any]:
- """Build the kwargs for the Post request, used by sync
-
- Args:
- prompt (str): prompt used in query
- kwargs (dict): model kwargs in payload
-
- Returns:
- Dict[str, Union[str,dict]]: _description_
- """
- _model_kwargs = self.model_kwargs or {}
- _params = {**_model_kwargs, **kwargs}
-
- return dict(
- url=f"{self.gradient_api_url}/models/{self.model_id}/complete",
- headers={
- "authorization": f"Bearer {self.gradient_access_token}",
- "x-gradient-workspace-id": f"{self.gradient_workspace_id}",
- "accept": "application/json",
- "content-type": "application/json",
- },
- json=dict(
- query=prompt,
- maxGeneratedTokenCount=_params.get("max_generated_token_count", None),
- temperature=_params.get("temperature", None),
- topK=_params.get("top_k", None),
- topP=_params.get("top_p", None),
- ),
- )
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call to Gradients API `model/{id}/complete`.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
- """
- try:
- response = requests.post(**self._kwargs_post_request(prompt, kwargs))
- if response.status_code != 200:
- raise Exception(
- f"Gradient returned an unexpected response with status "
- f"{response.status_code}: {response.text}"
- )
- except requests.exceptions.RequestException as e:
- raise Exception(f"RequestException while calling Gradient Endpoint: {e}")
-
- text = response.json()["generatedOutput"]
-
- if stop is not None:
- # Apply stop tokens when making calls to Gradient
- text = enforce_stop_tokens(text, stop)
-
- return text
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Async Call to Gradients API `model/{id}/complete`.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
- """
- if not self.aiosession:
- async with aiohttp.ClientSession() as session:
- async with session.post(
- **self._kwargs_post_request(prompt=prompt, kwargs=kwargs)
- ) as response:
- if response.status != 200:
- raise Exception(
- f"Gradient returned an unexpected response with status "
- f"{response.status}: {response.text}"
- )
- text = (await response.json())["generatedOutput"]
- else:
- async with self.aiosession.post(
- **self._kwargs_post_request(prompt=prompt, kwargs=kwargs)
- ) as response:
- if response.status != 200:
- raise Exception(
- f"Gradient returned an unexpected response with status "
- f"{response.status}: {response.text}"
- )
- text = (await response.json())["generatedOutput"]
-
- if stop is not None:
- # Apply stop tokens when making calls to Gradient
- text = enforce_stop_tokens(text, stop)
-
- return text
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Run the LLM on the given prompt and input."""
-
- # same thing with threading
- def _inner_generate(prompt: str) -> List[Generation]:
- return [
- Generation(
- text=self._call(
- prompt=prompt, stop=stop, run_manager=run_manager, **kwargs
- )
- )
- ]
-
- if len(prompts) <= 1:
- generations = list(map(_inner_generate, prompts))
- else:
- with ThreadPoolExecutor(min(8, len(prompts))) as p:
- generations = list(p.map(_inner_generate, prompts))
-
- return LLMResult(generations=generations)
-
- async def _agenerate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Run the LLM on the given prompt and input."""
- generations = []
- for generation in await asyncio.gather(
- *[
- self._acall(prompt, stop=stop, run_manager=run_manager, **kwargs)
- for prompt in prompts
- ]
- ):
- generations.append([Generation(text=generation)])
- return LLMResult(generations=generations)
-
- def train_unsupervised(
- self,
- inputs: Sequence[str],
- **kwargs: Any,
- ) -> TrainResult:
- try:
- response = requests.post(
- **self._kwargs_post_fine_tune_request(inputs, kwargs)
- )
- if response.status_code != 200:
- raise Exception(
- f"Gradient returned an unexpected response with status "
- f"{response.status_code}: {response.text}"
- )
- except requests.exceptions.RequestException as e:
- raise Exception(f"RequestException while calling Gradient Endpoint: {e}")
-
- response_json = response.json()
- loss = response_json["sumLoss"] / response_json["numberOfTrainableTokens"]
- return TrainResult(loss=loss)
-
- async def atrain_unsupervised(
- self,
- inputs: Sequence[str],
- **kwargs: Any,
- ) -> TrainResult:
- if not self.aiosession:
- async with aiohttp.ClientSession() as session:
- async with session.post(
- **self._kwargs_post_fine_tune_request(inputs, kwargs)
- ) as response:
- if response.status != 200:
- raise Exception(
- f"Gradient returned an unexpected response with status "
- f"{response.status}: {response.text}"
- )
- response_json = await response.json()
- loss = (
- response_json["sumLoss"]
- / response_json["numberOfTrainableTokens"]
- )
- else:
- async with self.aiosession.post(
- **self._kwargs_post_fine_tune_request(inputs, kwargs)
- ) as response:
- if response.status != 200:
- raise Exception(
- f"Gradient returned an unexpected response with status "
- f"{response.status}: {response.text}"
- )
- response_json = await response.json()
- loss = (
- response_json["sumLoss"] / response_json["numberOfTrainableTokens"]
- )
-
- return TrainResult(loss=loss)
diff --git a/libs/community/langchain_community/llms/grammars/json.gbnf b/libs/community/langchain_community/llms/grammars/json.gbnf
deleted file mode 100644
index 61bd2b2e65..0000000000
--- a/libs/community/langchain_community/llms/grammars/json.gbnf
+++ /dev/null
@@ -1,29 +0,0 @@
-# Grammar for subset of JSON - doesn't support full string or number syntax
-
-root ::= object
-value ::= object | array | string | number | boolean | "null"
-
-object ::=
- "{" ws (
- string ":" ws value
- ("," ws string ":" ws value)*
- )? "}"
-
-array ::=
- "[" ws (
- value
- ("," ws value)*
- )? "]"
-
-string ::=
- "\"" (
- [^"\\] |
- "\\" (["\\/bfnrt] | "u" [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F]) # escapes
- )* "\"" ws
-
-# Only plain integers currently
-number ::= "-"? [0-9]+ ws
-boolean ::= ("true" | "false") ws
-
-# Optional space: by convention, applied in this grammar after literal chars when allowed
-ws ::= ([ \t\n] ws)?
\ No newline at end of file
diff --git a/libs/community/langchain_community/llms/grammars/list.gbnf b/libs/community/langchain_community/llms/grammars/list.gbnf
deleted file mode 100644
index 30ea6e0c84..0000000000
--- a/libs/community/langchain_community/llms/grammars/list.gbnf
+++ /dev/null
@@ -1,14 +0,0 @@
-root ::= "[" items "]" EOF
-
-items ::= item ("," ws* item)*
-
-item ::= string
-
-string ::=
- "\"" word (ws+ word)* "\"" ws*
-
-word ::= [a-zA-Z]+
-
-ws ::= " "
-
-EOF ::= "\n"
\ No newline at end of file
diff --git a/libs/community/langchain_community/llms/huggingface_endpoint.py b/libs/community/langchain_community/llms/huggingface_endpoint.py
deleted file mode 100644
index 61efcc3f36..0000000000
--- a/libs/community/langchain_community/llms/huggingface_endpoint.py
+++ /dev/null
@@ -1,389 +0,0 @@
-import json
-import logging
-import os
-from typing import Any, AsyncIterator, Dict, Iterator, List, Mapping, Optional
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import (
- get_pydantic_field_names,
- pre_init,
-)
-from pydantic import ConfigDict, Field, model_validator
-
-logger = logging.getLogger(__name__)
-
-VALID_TASKS = (
- "text2text-generation",
- "text-generation",
- "summarization",
- "conversational",
-)
-
-
-@deprecated(
- since="0.0.37",
- removal="1.0",
- alternative_import="langchain_huggingface.HuggingFaceEndpoint",
-)
-class HuggingFaceEndpoint(LLM):
- """
- HuggingFace Endpoint.
-
- To use this class, you should have installed the ``huggingface_hub`` package, and
- the environment variable ``HUGGINGFACEHUB_API_TOKEN`` set with your API token,
- or given as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- # Basic Example (no streaming)
- llm = HuggingFaceEndpoint(
- endpoint_url="http://localhost:8010/",
- max_new_tokens=512,
- top_k=10,
- top_p=0.95,
- typical_p=0.95,
- temperature=0.01,
- repetition_penalty=1.03,
- huggingfacehub_api_token="my-api-key"
- )
- print(llm.invoke("What is Deep Learning?"))
-
- # Streaming response example
- from langchain_core.callbacks.streaming_stdout import StreamingStdOutCallbackHandler
-
- callbacks = [StreamingStdOutCallbackHandler()]
- llm = HuggingFaceEndpoint(
- endpoint_url="http://localhost:8010/",
- max_new_tokens=512,
- top_k=10,
- top_p=0.95,
- typical_p=0.95,
- temperature=0.01,
- repetition_penalty=1.03,
- callbacks=callbacks,
- streaming=True,
- huggingfacehub_api_token="my-api-key"
- )
- print(llm.invoke("What is Deep Learning?"))
-
- """ # noqa: E501
-
- endpoint_url: Optional[str] = None
- """Endpoint URL to use."""
- repo_id: Optional[str] = None
- """Repo to use."""
- huggingfacehub_api_token: Optional[str] = None
- max_new_tokens: int = 512
- """Maximum number of generated tokens"""
- top_k: Optional[int] = None
- """The number of highest probability vocabulary tokens to keep for
- top-k-filtering."""
- top_p: Optional[float] = 0.95
- """If set to < 1, only the smallest set of most probable tokens with probabilities
- that add up to `top_p` or higher are kept for generation."""
- typical_p: Optional[float] = 0.95
- """Typical Decoding mass. See [Typical Decoding for Natural Language
- Generation](https://arxiv.org/abs/2202.00666) for more information."""
- temperature: Optional[float] = 0.8
- """The value used to module the logits distribution."""
- repetition_penalty: Optional[float] = None
- """The parameter for repetition penalty. 1.0 means no penalty.
- See [this paper](https://arxiv.org/pdf/1909.05858.pdf) for more details."""
- return_full_text: bool = False
- """Whether to prepend the prompt to the generated text"""
- truncate: Optional[int] = None
- """Truncate inputs tokens to the given size"""
- stop_sequences: List[str] = Field(default_factory=list)
- """Stop generating tokens if a member of `stop_sequences` is generated"""
- seed: Optional[int] = None
- """Random sampling seed"""
- inference_server_url: str = ""
- """text-generation-inference instance base url"""
- timeout: int = 120
- """Timeout in seconds"""
- streaming: bool = False
- """Whether to generate a stream of tokens asynchronously"""
- do_sample: bool = False
- """Activate logits sampling"""
- watermark: bool = False
- """Watermarking with [A Watermark for Large Language Models]
- (https://arxiv.org/abs/2301.10226)"""
- server_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any text-generation-inference server parameters not explicitly specified"""
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `call` not explicitly specified"""
- model: str
- client: Any = None
- async_client: Any = None
- task: Optional[str] = None
- """Task to call the model with.
- Should be a task that returns `generated_text` or `summary_text`."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = get_pydantic_field_names(cls)
- extra = values.get("model_kwargs", {})
- for field_name in list(values):
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- if field_name not in all_required_field_names:
- logger.warning(
- f"""WARNING! {field_name} is not default parameter.
- {field_name} was transferred to model_kwargs.
- Please make sure that {field_name} is what you intended."""
- )
- extra[field_name] = values.pop(field_name)
-
- invalid_model_kwargs = all_required_field_names.intersection(extra.keys())
- if invalid_model_kwargs:
- raise ValueError(
- f"Parameters {invalid_model_kwargs} should be specified explicitly. "
- f"Instead they were passed in as part of `model_kwargs` parameter."
- )
-
- values["model_kwargs"] = extra
- if "endpoint_url" not in values and "repo_id" not in values:
- raise ValueError(
- "Please specify an `endpoint_url` or `repo_id` for the model."
- )
- if "endpoint_url" in values and "repo_id" in values:
- raise ValueError(
- "Please specify either an `endpoint_url` OR a `repo_id`, not both."
- )
- values["model"] = values.get("endpoint_url") or values.get("repo_id")
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that package is installed and that the API token is valid."""
- try:
- from huggingface_hub import login
-
- except ImportError:
- raise ImportError(
- "Could not import huggingface_hub python package. "
- "Please install it with `pip install huggingface_hub`."
- )
- huggingfacehub_api_token = values["huggingfacehub_api_token"] or os.getenv(
- "HUGGINGFACEHUB_API_TOKEN"
- )
- if huggingfacehub_api_token is not None:
- try:
- login(token=huggingfacehub_api_token)
- except Exception as e:
- raise ValueError(
- "Could not authenticate with huggingface_hub. "
- "Please check your API token."
- ) from e
-
- from huggingface_hub import AsyncInferenceClient, InferenceClient
-
- values["client"] = InferenceClient(
- model=values["model"],
- timeout=values["timeout"],
- token=huggingfacehub_api_token,
- **values["server_kwargs"],
- )
- values["async_client"] = AsyncInferenceClient(
- model=values["model"],
- timeout=values["timeout"],
- token=huggingfacehub_api_token,
- **values["server_kwargs"],
- )
-
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling text generation inference API."""
- return {
- "max_new_tokens": self.max_new_tokens,
- "top_k": self.top_k,
- "top_p": self.top_p,
- "typical_p": self.typical_p,
- "temperature": self.temperature,
- "repetition_penalty": self.repetition_penalty,
- "return_full_text": self.return_full_text,
- "truncate": self.truncate,
- "stop_sequences": self.stop_sequences,
- "seed": self.seed,
- "do_sample": self.do_sample,
- "watermark": self.watermark,
- **self.model_kwargs,
- }
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- **{"endpoint_url": self.endpoint_url, "task": self.task},
- **{"model_kwargs": _model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "huggingface_endpoint"
-
- def _invocation_params(
- self, runtime_stop: Optional[List[str]], **kwargs: Any
- ) -> Dict[str, Any]:
- params = {**self._default_params, **kwargs}
- params["stop_sequences"] = params["stop_sequences"] + (runtime_stop or [])
- return params
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to HuggingFace Hub's inference endpoint."""
- invocation_params = self._invocation_params(stop, **kwargs)
- if self.streaming:
- completion = ""
- for chunk in self._stream(prompt, stop, run_manager, **invocation_params):
- completion += chunk.text
- return completion
- else:
- invocation_params["stop"] = invocation_params[
- "stop_sequences"
- ] # porting 'stop_sequences' into the 'stop' argument
- response = self.client.post(
- json={"inputs": prompt, "parameters": invocation_params},
- stream=False,
- task=self.task,
- )
- try:
- response_text = json.loads(response.decode())[0]["generated_text"]
- except KeyError:
- response_text = json.loads(response.decode())["generated_text"]
-
- # Maybe the generation has stopped at one of the stop sequences:
- # then we remove this stop sequence from the end of the generated text
- for stop_seq in invocation_params["stop_sequences"]:
- if response_text[-len(stop_seq) :] == stop_seq:
- response_text = response_text[: -len(stop_seq)]
- return response_text
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- invocation_params = self._invocation_params(stop, **kwargs)
- if self.streaming:
- completion = ""
- async for chunk in self._astream(
- prompt, stop, run_manager, **invocation_params
- ):
- completion += chunk.text
- return completion
- else:
- invocation_params["stop"] = invocation_params["stop_sequences"]
- response = await self.async_client.post(
- json={"inputs": prompt, "parameters": invocation_params},
- stream=False,
- task=self.task,
- )
- try:
- response_text = json.loads(response.decode())[0]["generated_text"]
- except KeyError:
- response_text = json.loads(response.decode())["generated_text"]
-
- # Maybe the generation has stopped at one of the stop sequences:
- # then remove this stop sequence from the end of the generated text
- for stop_seq in invocation_params["stop_sequences"]:
- if response_text[-len(stop_seq) :] == stop_seq:
- response_text = response_text[: -len(stop_seq)]
- return response_text
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- invocation_params = self._invocation_params(stop, **kwargs)
-
- for response in self.client.text_generation(
- prompt, **invocation_params, stream=True
- ):
- # identify stop sequence in generated text, if any
- stop_seq_found: Optional[str] = None
- for stop_seq in invocation_params["stop_sequences"]:
- if stop_seq in response:
- stop_seq_found = stop_seq
-
- # identify text to yield
- text: Optional[str] = None
- if stop_seq_found:
- text = response[: response.index(stop_seq_found)]
- else:
- text = response
-
- # yield text, if any
- if text:
- chunk = GenerationChunk(text=text)
-
- if run_manager:
- run_manager.on_llm_new_token(chunk.text)
- yield chunk
-
- # break if stop sequence found
- if stop_seq_found:
- break
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- invocation_params = self._invocation_params(stop, **kwargs)
- async for response in await self.async_client.text_generation(
- prompt, **invocation_params, stream=True
- ):
- # identify stop sequence in generated text, if any
- stop_seq_found: Optional[str] = None
- for stop_seq in invocation_params["stop_sequences"]:
- if stop_seq in response:
- stop_seq_found = stop_seq
-
- # identify text to yield
- text: Optional[str] = None
- if stop_seq_found:
- text = response[: response.index(stop_seq_found)]
- else:
- text = response
-
- # yield text, if any
- if text:
- chunk = GenerationChunk(text=text)
-
- if run_manager:
- await run_manager.on_llm_new_token(chunk.text)
- yield chunk
-
- # break if stop sequence found
- if stop_seq_found:
- break
diff --git a/libs/community/langchain_community/llms/huggingface_hub.py b/libs/community/langchain_community/llms/huggingface_hub.py
deleted file mode 100644
index 95e8d8f0d2..0000000000
--- a/libs/community/langchain_community/llms/huggingface_hub.py
+++ /dev/null
@@ -1,155 +0,0 @@
-import json
-from typing import Any, Dict, List, Mapping, Optional
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import ConfigDict
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-# key: task
-# value: key in the output dictionary
-VALID_TASKS_DICT = {
- "translation": "translation_text",
- "summarization": "summary_text",
- "conversational": "generated_text",
- "text-generation": "generated_text",
- "text2text-generation": "generated_text",
-}
-
-
-@deprecated(
- "0.0.21",
- removal="1.0",
- alternative_import="langchain_huggingface.HuggingFaceEndpoint",
-)
-class HuggingFaceHub(LLM):
- """HuggingFaceHub models.
- ! This class is deprecated, you should use HuggingFaceEndpoint instead.
-
- To use, you should have the ``huggingface_hub`` python package installed, and the
- environment variable ``HUGGINGFACEHUB_API_TOKEN`` set with your API token, or pass
- it as a named parameter to the constructor.
-
- Supports `text-generation`, `text2text-generation`, `conversational`, `translation`,
- and `summarization`.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import HuggingFaceHub
- hf = HuggingFaceHub(repo_id="gpt2", huggingfacehub_api_token="my-api-key")
- """
-
- client: Any = None #: :meta private:
- repo_id: Optional[str] = None
- """Model name to use.
- If not provided, the default model for the chosen task will be used."""
- task: Optional[str] = None
- """Task to call the model with.
- Should be a task that returns `generated_text`, `summary_text`,
- or `translation_text`."""
- model_kwargs: Optional[dict] = None
- """Keyword arguments to pass to the model."""
-
- huggingfacehub_api_token: Optional[str] = None
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- huggingfacehub_api_token = get_from_dict_or_env(
- values, "huggingfacehub_api_token", "HUGGINGFACEHUB_API_TOKEN"
- )
- try:
- from huggingface_hub import HfApi, InferenceClient
-
- repo_id = values["repo_id"]
- client = InferenceClient(
- model=repo_id,
- token=huggingfacehub_api_token,
- )
- if not values["task"]:
- if not repo_id:
- raise ValueError(
- "Must specify either `repo_id` or `task`, or both."
- )
- # Use the recommended task for the chosen model
- model_info = HfApi(token=huggingfacehub_api_token).model_info(
- repo_id=repo_id
- )
- values["task"] = model_info.pipeline_tag
- if values["task"] not in VALID_TASKS_DICT:
- raise ValueError(
- f"Got invalid task {values['task']}, "
- f"currently only {VALID_TASKS_DICT.keys()} are supported"
- )
- values["client"] = client
- except ImportError:
- raise ImportError(
- "Could not import huggingface_hub python package. "
- "Please install it with `pip install huggingface_hub`."
- )
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- **{"repo_id": self.repo_id, "task": self.task},
- **{"model_kwargs": _model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "huggingface_hub"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to HuggingFace Hub's inference endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = hf("Tell me a joke.")
- """
- _model_kwargs = self.model_kwargs or {}
- parameters = {**_model_kwargs, **kwargs}
-
- response = self.client.post(
- json={"inputs": prompt, "parameters": parameters}, task=self.task
- )
- response = json.loads(response.decode())
- if "error" in response:
- raise ValueError(f"Error raised by inference API: {response['error']}")
-
- response_key = VALID_TASKS_DICT[self.task] # type: ignore[index]
- if isinstance(response, list):
- text = response[0][response_key]
- else:
- text = response[response_key]
-
- if stop is not None:
- # This is a bit hacky, but I can't figure out a better way to enforce
- # stop tokens when making calls to huggingface_hub.
- text = enforce_stop_tokens(text, stop)
- return text
diff --git a/libs/community/langchain_community/llms/huggingface_pipeline.py b/libs/community/langchain_community/llms/huggingface_pipeline.py
deleted file mode 100644
index 185405645e..0000000000
--- a/libs/community/langchain_community/llms/huggingface_pipeline.py
+++ /dev/null
@@ -1,376 +0,0 @@
-from __future__ import annotations
-
-import importlib.util
-import logging
-from typing import Any, Iterator, List, Mapping, Optional
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import BaseLLM
-from langchain_core.outputs import Generation, GenerationChunk, LLMResult
-from pydantic import ConfigDict
-
-DEFAULT_MODEL_ID = "gpt2"
-DEFAULT_TASK = "text-generation"
-VALID_TASKS = (
- "text2text-generation",
- "text-generation",
- "summarization",
- "translation",
-)
-DEFAULT_BATCH_SIZE = 4
-
-logger = logging.getLogger(__name__)
-
-
-@deprecated(
- since="0.0.37",
- removal="1.0",
- alternative_import="langchain_huggingface.HuggingFacePipeline",
-)
-class HuggingFacePipeline(BaseLLM):
- """HuggingFace Pipeline API.
-
- To use, you should have the ``transformers`` python package installed.
-
- Only supports `text-generation`, `text2text-generation`, `summarization` and
- `translation` for now.
-
- Example using from_model_id:
- .. code-block:: python
-
- from langchain_community.llms import HuggingFacePipeline
- hf = HuggingFacePipeline.from_model_id(
- model_id="gpt2",
- task="text-generation",
- pipeline_kwargs={"max_new_tokens": 10},
- )
- Example passing pipeline in directly:
- .. code-block:: python
-
- from langchain_community.llms import HuggingFacePipeline
- from transformers import AutoModelForCausalLM, AutoTokenizer, pipeline
-
- model_id = "gpt2"
- tokenizer = AutoTokenizer.from_pretrained(model_id)
- model = AutoModelForCausalLM.from_pretrained(model_id)
- pipe = pipeline(
- "text-generation", model=model, tokenizer=tokenizer, max_new_tokens=10
- )
- hf = HuggingFacePipeline(pipeline=pipe)
- """
-
- pipeline: Any = None #: :meta private:
- model_id: str = DEFAULT_MODEL_ID
- """Model name to use."""
- model_kwargs: Optional[dict] = None
- """Keyword arguments passed to the model."""
- pipeline_kwargs: Optional[dict] = None
- """Keyword arguments passed to the pipeline."""
- batch_size: int = DEFAULT_BATCH_SIZE
- """Batch size to use when passing multiple documents to generate."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @classmethod
- def from_model_id(
- cls,
- model_id: str,
- task: str,
- backend: str = "default",
- device: Optional[int] = -1,
- device_map: Optional[str] = None,
- model_kwargs: Optional[dict] = None,
- pipeline_kwargs: Optional[dict] = None,
- batch_size: int = DEFAULT_BATCH_SIZE,
- **kwargs: Any,
- ) -> HuggingFacePipeline:
- """Construct the pipeline object from model_id and task."""
- try:
- from transformers import (
- AutoModelForCausalLM,
- AutoModelForSeq2SeqLM,
- AutoTokenizer,
- )
- from transformers import pipeline as hf_pipeline
-
- except ImportError:
- raise ImportError(
- "Could not import transformers python package. "
- "Please install it with `pip install transformers`."
- )
-
- _model_kwargs = model_kwargs or {}
- tokenizer = AutoTokenizer.from_pretrained(model_id, **_model_kwargs)
-
- try:
- if task == "text-generation":
- if backend == "openvino":
- try:
- from optimum.intel.openvino import OVModelForCausalLM
-
- except ImportError:
- raise ImportError(
- "Could not import optimum-intel python package. "
- "Please install it with: "
- "pip install 'optimum[openvino,nncf]' "
- )
- try:
- # use local model
- model = OVModelForCausalLM.from_pretrained(
- model_id, **_model_kwargs
- )
-
- except Exception:
- # use remote model
- model = OVModelForCausalLM.from_pretrained(
- model_id, export=True, **_model_kwargs
- )
- else:
- model = AutoModelForCausalLM.from_pretrained(
- model_id, **_model_kwargs
- )
- elif task in ("text2text-generation", "summarization", "translation"):
- if backend == "openvino":
- try:
- from optimum.intel.openvino import OVModelForSeq2SeqLM
-
- except ImportError:
- raise ImportError(
- "Could not import optimum-intel python package. "
- "Please install it with: "
- "pip install 'optimum[openvino,nncf]' "
- )
- try:
- # use local model
- model = OVModelForSeq2SeqLM.from_pretrained(
- model_id, **_model_kwargs
- )
-
- except Exception:
- # use remote model
- model = OVModelForSeq2SeqLM.from_pretrained(
- model_id, export=True, **_model_kwargs
- )
- else:
- model = AutoModelForSeq2SeqLM.from_pretrained(
- model_id, **_model_kwargs
- )
- else:
- raise ValueError(
- f"Got invalid task {task}, "
- f"currently only {VALID_TASKS} are supported"
- )
- except ImportError as e:
- raise ImportError(
- f"Could not load the {task} model due to missing dependencies."
- ) from e
-
- if tokenizer.pad_token is None:
- if model.config.pad_token_id is not None:
- tokenizer.pad_token_id = model.config.pad_token_id
- elif model.config.eos_token_id is not None and isinstance(
- model.config.eos_token_id, int
- ):
- tokenizer.pad_token_id = model.config.eos_token_id
- elif tokenizer.eos_token_id is not None:
- tokenizer.pad_token_id = tokenizer.eos_token_id
- else:
- tokenizer.add_special_tokens({"pad_token": "[PAD]"})
-
- if (
- (
- getattr(model, "is_loaded_in_4bit", False)
- or getattr(model, "is_loaded_in_8bit", False)
- )
- and device is not None
- and backend == "default"
- ):
- logger.warning(
- f"Setting the `device` argument to None from {device} to avoid "
- "the error caused by attempting to move the model that was already "
- "loaded on the GPU using the Accelerate module to the same or "
- "another device."
- )
- device = None
-
- if (
- device is not None
- and importlib.util.find_spec("torch") is not None
- and backend == "default"
- ):
- import torch
-
- cuda_device_count = torch.cuda.device_count()
- if device < -1 or (device >= cuda_device_count):
- raise ValueError(
- f"Got device=={device}, "
- f"device is required to be within [-1, {cuda_device_count})"
- )
- if device_map is not None and device < 0:
- device = None
- if device is not None and device < 0 and cuda_device_count > 0:
- logger.warning(
- "Device has %d GPUs available. "
- "Provide device={deviceId} to `from_model_id` to use available"
- "GPUs for execution. deviceId is -1 (default) for CPU and "
- "can be a positive integer associated with CUDA device id.",
- cuda_device_count,
- )
- if device is not None and device_map is not None and backend == "openvino":
- logger.warning("Please set device for OpenVINO through: `model_kwargs`")
- if "trust_remote_code" in _model_kwargs:
- _model_kwargs = {
- k: v for k, v in _model_kwargs.items() if k != "trust_remote_code"
- }
- _pipeline_kwargs = pipeline_kwargs or {}
- pipeline = hf_pipeline(
- task=task,
- model=model,
- tokenizer=tokenizer,
- device=device,
- device_map=device_map,
- batch_size=batch_size,
- model_kwargs=_model_kwargs,
- **_pipeline_kwargs,
- )
- if pipeline.task not in VALID_TASKS:
- raise ValueError(
- f"Got invalid task {pipeline.task}, "
- f"currently only {VALID_TASKS} are supported"
- )
- return cls(
- pipeline=pipeline,
- model_id=model_id,
- model_kwargs=_model_kwargs,
- pipeline_kwargs=_pipeline_kwargs,
- batch_size=batch_size,
- **kwargs,
- )
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- "model_id": self.model_id,
- "model_kwargs": self.model_kwargs,
- "pipeline_kwargs": self.pipeline_kwargs,
- }
-
- @property
- def _llm_type(self) -> str:
- return "huggingface_pipeline"
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- # List to hold all results
- text_generations: List[str] = []
-
- default_pipeline_kwargs = self.pipeline_kwargs if self.pipeline_kwargs else {}
- pipeline_kwargs = kwargs.get("pipeline_kwargs", default_pipeline_kwargs)
-
- skip_prompt = kwargs.get("skip_prompt", False)
-
- for i in range(0, len(prompts), self.batch_size):
- batch_prompts = prompts[i : i + self.batch_size]
-
- # Process batch of prompts
- responses = self.pipeline(
- batch_prompts,
- **pipeline_kwargs,
- )
-
- # Process each response in the batch
- for j, response in enumerate(responses):
- if isinstance(response, list):
- # if model returns multiple generations, pick the top one
- response = response[0]
-
- if self.pipeline.task == "text-generation":
- text = response["generated_text"]
- elif self.pipeline.task == "text2text-generation":
- text = response["generated_text"]
- elif self.pipeline.task == "summarization":
- text = response["summary_text"]
- elif self.pipeline.task in "translation":
- text = response["translation_text"]
- else:
- raise ValueError(
- f"Got invalid task {self.pipeline.task}, "
- f"currently only {VALID_TASKS} are supported"
- )
- if skip_prompt:
- text = text[len(batch_prompts[j]) :]
- # Append the processed text to results
- text_generations.append(text)
-
- return LLMResult(
- generations=[[Generation(text=text)] for text in text_generations]
- )
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- from threading import Thread
-
- import torch
- from transformers import (
- StoppingCriteria,
- StoppingCriteriaList,
- TextIteratorStreamer,
- )
-
- pipeline_kwargs = kwargs.get("pipeline_kwargs", {})
- skip_prompt = kwargs.get("skip_prompt", True)
-
- if stop is not None:
- stop = self.pipeline.tokenizer.convert_tokens_to_ids(stop)
- stopping_ids_list = stop or []
-
- class StopOnTokens(StoppingCriteria):
- def __call__(
- self,
- input_ids: torch.LongTensor,
- scores: torch.FloatTensor,
- **kwargs: Any,
- ) -> bool:
- for stop_id in stopping_ids_list:
- if input_ids[0][-1] == stop_id:
- return True
- return False
-
- stopping_criteria = StoppingCriteriaList([StopOnTokens()])
-
- inputs = self.pipeline.tokenizer(prompt, return_tensors="pt")
- streamer = TextIteratorStreamer(
- self.pipeline.tokenizer,
- timeout=60.0,
- skip_prompt=skip_prompt,
- skip_special_tokens=True,
- )
- generation_kwargs = dict(
- inputs,
- streamer=streamer,
- stopping_criteria=stopping_criteria,
- **pipeline_kwargs,
- )
- t1 = Thread(target=self.pipeline.model.generate, kwargs=generation_kwargs)
- t1.start()
-
- for char in streamer:
- chunk = GenerationChunk(text=char)
- if run_manager:
- run_manager.on_llm_new_token(chunk.text, chunk=chunk)
-
- yield chunk
diff --git a/libs/community/langchain_community/llms/huggingface_text_gen_inference.py b/libs/community/langchain_community/llms/huggingface_text_gen_inference.py
deleted file mode 100644
index e536a2ee4a..0000000000
--- a/libs/community/langchain_community/llms/huggingface_text_gen_inference.py
+++ /dev/null
@@ -1,310 +0,0 @@
-import logging
-from typing import Any, AsyncIterator, Dict, Iterator, List, Optional
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import get_pydantic_field_names, pre_init
-from pydantic import ConfigDict, Field, model_validator
-
-logger = logging.getLogger(__name__)
-
-
-@deprecated(
- "0.0.21",
- removal="1.0",
- alternative_import="langchain_huggingface.HuggingFaceEndpoint",
-)
-class HuggingFaceTextGenInference(LLM):
- """
- HuggingFace text generation API.
- ! This class is deprecated, you should use HuggingFaceEndpoint instead !
-
- To use, you should have the `text-generation` python package installed and
- a text-generation server running.
-
- Example:
- .. code-block:: python
-
- # Basic Example (no streaming)
- llm = HuggingFaceTextGenInference(
- inference_server_url="http://localhost:8010/",
- max_new_tokens=512,
- top_k=10,
- top_p=0.95,
- typical_p=0.95,
- temperature=0.01,
- repetition_penalty=1.03,
- )
- print(llm.invoke("What is Deep Learning?")) # noqa: T201
-
- # Streaming response example
- from langchain_community.callbacks import streaming_stdout
-
- callbacks = [streaming_stdout.StreamingStdOutCallbackHandler()]
- llm = HuggingFaceTextGenInference(
- inference_server_url="http://localhost:8010/",
- max_new_tokens=512,
- top_k=10,
- top_p=0.95,
- typical_p=0.95,
- temperature=0.01,
- repetition_penalty=1.03,
- callbacks=callbacks,
- streaming=True
- )
- print(llm.invoke("What is Deep Learning?")) # noqa: T201
-
- """
-
- max_new_tokens: int = 512
- """Maximum number of generated tokens"""
- top_k: Optional[int] = None
- """The number of highest probability vocabulary tokens to keep for
- top-k-filtering."""
- top_p: Optional[float] = 0.95
- """If set to < 1, only the smallest set of most probable tokens with probabilities
- that add up to `top_p` or higher are kept for generation."""
- typical_p: Optional[float] = 0.95
- """Typical Decoding mass. See [Typical Decoding for Natural Language
- Generation](https://arxiv.org/abs/2202.00666) for more information."""
- temperature: Optional[float] = 0.8
- """The value used to module the logits distribution."""
- repetition_penalty: Optional[float] = None
- """The parameter for repetition penalty. 1.0 means no penalty.
- See [this paper](https://arxiv.org/pdf/1909.05858.pdf) for more details."""
- return_full_text: bool = False
- """Whether to prepend the prompt to the generated text"""
- truncate: Optional[int] = None
- """Truncate inputs tokens to the given size"""
- stop_sequences: List[str] = Field(default_factory=list)
- """Stop generating tokens if a member of `stop_sequences` is generated"""
- seed: Optional[int] = None
- """Random sampling seed"""
- inference_server_url: str = ""
- """text-generation-inference instance base url"""
- timeout: int = 120
- """Timeout in seconds"""
- streaming: bool = False
- """Whether to generate a stream of tokens asynchronously"""
- do_sample: bool = False
- """Activate logits sampling"""
- watermark: bool = False
- """Watermarking with [A Watermark for Large Language Models]
- (https://arxiv.org/abs/2301.10226)"""
- server_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any text-generation-inference server parameters not explicitly specified"""
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `call` not explicitly specified"""
- client: Any = None
- async_client: Any = None
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = get_pydantic_field_names(cls)
- extra = values.get("model_kwargs", {})
- for field_name in list(values):
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- if field_name not in all_required_field_names:
- logger.warning(
- f"""WARNING! {field_name} is not default parameter.
- {field_name} was transferred to model_kwargs.
- Please confirm that {field_name} is what you intended."""
- )
- extra[field_name] = values.pop(field_name)
-
- invalid_model_kwargs = all_required_field_names.intersection(extra.keys())
- if invalid_model_kwargs:
- raise ValueError(
- f"Parameters {invalid_model_kwargs} should be specified explicitly. "
- f"Instead they were passed in as part of `model_kwargs` parameter."
- )
-
- values["model_kwargs"] = extra
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that python package exists in environment."""
-
- try:
- import text_generation
-
- values["client"] = text_generation.Client(
- values["inference_server_url"],
- timeout=values["timeout"],
- **values["server_kwargs"],
- )
- values["async_client"] = text_generation.AsyncClient(
- values["inference_server_url"],
- timeout=values["timeout"],
- **values["server_kwargs"],
- )
- except ImportError:
- raise ImportError(
- "Could not import text_generation python package. "
- "Please install it with `pip install text_generation`."
- )
- return values
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "huggingface_textgen_inference"
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling text generation inference API."""
- return {
- "max_new_tokens": self.max_new_tokens,
- "top_k": self.top_k,
- "top_p": self.top_p,
- "typical_p": self.typical_p,
- "temperature": self.temperature,
- "repetition_penalty": self.repetition_penalty,
- "return_full_text": self.return_full_text,
- "truncate": self.truncate,
- "stop_sequences": self.stop_sequences,
- "seed": self.seed,
- "do_sample": self.do_sample,
- "watermark": self.watermark,
- **self.model_kwargs,
- }
-
- def _invocation_params(
- self, runtime_stop: Optional[List[str]], **kwargs: Any
- ) -> Dict[str, Any]:
- params = {**self._default_params, **kwargs}
- params["stop_sequences"] = params["stop_sequences"] + (runtime_stop or [])
- return params
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- if self.streaming:
- completion = ""
- for chunk in self._stream(prompt, stop, run_manager, **kwargs):
- completion += chunk.text
- return completion
-
- invocation_params = self._invocation_params(stop, **kwargs)
- res = self.client.generate(prompt, **invocation_params)
- # remove stop sequences from the end of the generated text
- for stop_seq in invocation_params["stop_sequences"]:
- if stop_seq in res.generated_text:
- res.generated_text = res.generated_text[
- : res.generated_text.index(stop_seq)
- ]
- return res.generated_text
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- if self.streaming:
- completion = ""
- async for chunk in self._astream(prompt, stop, run_manager, **kwargs):
- completion += chunk.text
- return completion
-
- invocation_params = self._invocation_params(stop, **kwargs)
- res = await self.async_client.generate(prompt, **invocation_params)
- # remove stop sequences from the end of the generated text
- for stop_seq in invocation_params["stop_sequences"]:
- if stop_seq in res.generated_text:
- res.generated_text = res.generated_text[
- : res.generated_text.index(stop_seq)
- ]
- return res.generated_text
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- invocation_params = self._invocation_params(stop, **kwargs)
-
- for res in self.client.generate_stream(prompt, **invocation_params):
- # identify stop sequence in generated text, if any
- stop_seq_found: Optional[str] = None
- for stop_seq in invocation_params["stop_sequences"]:
- if stop_seq in res.token.text:
- stop_seq_found = stop_seq
-
- # identify text to yield
- text: Optional[str] = None
- if res.token.special:
- text = None
- elif stop_seq_found:
- text = res.token.text[: res.token.text.index(stop_seq_found)]
- else:
- text = res.token.text
-
- # yield text, if any
- if text:
- chunk = GenerationChunk(text=text)
-
- if run_manager:
- run_manager.on_llm_new_token(chunk.text)
- yield chunk
-
- # break if stop sequence found
- if stop_seq_found:
- break
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- invocation_params = self._invocation_params(stop, **kwargs)
-
- async for res in self.async_client.generate_stream(prompt, **invocation_params):
- # identify stop sequence in generated text, if any
- stop_seq_found: Optional[str] = None
- for stop_seq in invocation_params["stop_sequences"]:
- if stop_seq in res.token.text:
- stop_seq_found = stop_seq
-
- # identify text to yield
- text: Optional[str] = None
- if res.token.special:
- text = None
- elif stop_seq_found:
- text = res.token.text[: res.token.text.index(stop_seq_found)]
- else:
- text = res.token.text
-
- # yield text, if any
- if text:
- chunk = GenerationChunk(text=text)
-
- if run_manager:
- await run_manager.on_llm_new_token(chunk.text)
- yield chunk
-
- # break if stop sequence found
- if stop_seq_found:
- break
diff --git a/libs/community/langchain_community/llms/human.py b/libs/community/langchain_community/llms/human.py
deleted file mode 100644
index 9a54b29aa1..0000000000
--- a/libs/community/langchain_community/llms/human.py
+++ /dev/null
@@ -1,83 +0,0 @@
-from typing import Any, Callable, List, Mapping, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from pydantic import Field
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-
-def _display_prompt(prompt: str) -> None:
- """Displays the given prompt to the user."""
- print(f"\n{prompt}") # noqa: T201
-
-
-def _collect_user_input(
- separator: Optional[str] = None, stop: Optional[List[str]] = None
-) -> str:
- """Collects and returns user input as a single string."""
- separator = separator or "\n"
- lines = []
-
- while True:
- line = input()
- if not line:
- break
- lines.append(line)
-
- if stop and any(seq in line for seq in stop):
- break
- # Combine all lines into a single string
- multi_line_input = separator.join(lines)
- return multi_line_input
-
-
-class HumanInputLLM(LLM):
- """User input as the response."""
-
- input_func: Callable = Field(default_factory=lambda: _collect_user_input)
- prompt_func: Callable[[str], None] = Field(default_factory=lambda: _display_prompt)
- separator: str = "\n"
- input_kwargs: Mapping[str, Any] = {}
- prompt_kwargs: Mapping[str, Any] = {}
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """
- Returns an empty dictionary as there are no identifying parameters.
- """
- return {}
-
- @property
- def _llm_type(self) -> str:
- """Returns the type of LLM."""
- return "human-input"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """
- Displays the prompt to the user and returns their input as a response.
-
- Args:
- prompt (str): The prompt to be displayed to the user.
- stop (Optional[List[str]]): A list of stop strings.
- run_manager (Optional[CallbackManagerForLLMRun]): Currently not used.
-
- Returns:
- str: The user's input as a response.
- """
- self.prompt_func(prompt, **self.prompt_kwargs)
- user_input = self.input_func(
- separator=self.separator, stop=stop, **self.input_kwargs
- )
-
- if stop is not None:
- # I believe this is required since the stop tokens
- # are not enforced by the human themselves
- user_input = enforce_stop_tokens(user_input, stop)
- return user_input
diff --git a/libs/community/langchain_community/llms/ipex_llm.py b/libs/community/langchain_community/llms/ipex_llm.py
deleted file mode 100644
index 0432b6aecc..0000000000
--- a/libs/community/langchain_community/llms/ipex_llm.py
+++ /dev/null
@@ -1,297 +0,0 @@
-import logging
-from typing import Any, List, Mapping, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from pydantic import ConfigDict
-
-DEFAULT_MODEL_ID = "gpt2"
-
-
-logger = logging.getLogger(__name__)
-
-
-class IpexLLM(LLM):
- """IpexLLM model.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import IpexLLM
- llm = IpexLLM.from_model_id(model_id="THUDM/chatglm-6b")
- """
-
- model_id: str = DEFAULT_MODEL_ID
- """Model name or model path to use."""
- model_kwargs: Optional[dict] = None
- """Keyword arguments passed to the model."""
- model: Any = None #: :meta private:
- """IpexLLM model."""
- tokenizer: Any = None #: :meta private:
- """Huggingface tokenizer model."""
- streaming: bool = True
- """Whether to stream the results, token by token."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @classmethod
- def from_model_id(
- cls,
- model_id: str,
- model_kwargs: Optional[dict] = None,
- *,
- tokenizer_id: Optional[str] = None,
- load_in_4bit: bool = True,
- load_in_low_bit: Optional[str] = None,
- **kwargs: Any,
- ) -> LLM:
- """
- Construct object from model_id
-
- Args:
- model_id: Path for the huggingface repo id to be downloaded or
- the huggingface checkpoint folder.
- tokenizer_id: Path for the huggingface repo id to be downloaded or
- the huggingface checkpoint folder which contains the tokenizer.
- load_in_4bit: "Whether to load model in 4bit.
- Unused if `load_in_low_bit` is not None.
- load_in_low_bit: Which low bit precisions to use when loading model.
- Example values: 'sym_int4', 'asym_int4', 'fp4', 'nf4', 'fp8', etc.
- Overrides `load_in_4bit` if specified.
- model_kwargs: Keyword arguments to pass to the model and tokenizer.
- kwargs: Extra arguments to pass to the model and tokenizer.
-
- Returns:
- An object of IpexLLM.
-
- """
-
- return cls._load_model(
- model_id=model_id,
- tokenizer_id=tokenizer_id,
- low_bit_model=False,
- load_in_4bit=load_in_4bit,
- load_in_low_bit=load_in_low_bit,
- model_kwargs=model_kwargs,
- kwargs=kwargs,
- )
-
- @classmethod
- def from_model_id_low_bit(
- cls,
- model_id: str,
- model_kwargs: Optional[dict] = None,
- *,
- tokenizer_id: Optional[str] = None,
- **kwargs: Any,
- ) -> LLM:
- """
- Construct low_bit object from model_id
-
- Args:
-
- model_id: Path for the ipex-llm transformers low-bit model folder.
- tokenizer_id: Path for the huggingface repo id or local model folder
- which contains the tokenizer.
- model_kwargs: Keyword arguments to pass to the model and tokenizer.
- kwargs: Extra arguments to pass to the model and tokenizer.
-
- Returns:
- An object of IpexLLM.
- """
-
- return cls._load_model(
- model_id=model_id,
- tokenizer_id=tokenizer_id,
- low_bit_model=True,
- load_in_4bit=False, # not used for low-bit model
- load_in_low_bit=None, # not used for low-bit model
- model_kwargs=model_kwargs,
- kwargs=kwargs,
- )
-
- @classmethod
- def _load_model(
- cls,
- model_id: str,
- tokenizer_id: Optional[str] = None,
- load_in_4bit: bool = False,
- load_in_low_bit: Optional[str] = None,
- low_bit_model: bool = False,
- model_kwargs: Optional[dict] = None,
- kwargs: Optional[dict] = None,
- ) -> Any:
- try:
- from ipex_llm.transformers import (
- AutoModel,
- AutoModelForCausalLM,
- )
- from transformers import AutoTokenizer, LlamaTokenizer
-
- except ImportError:
- raise ImportError(
- "Could not import ipex-llm. "
- "Please install `ipex-llm` properly following installation guides: "
- "https://github.com/intel-analytics/ipex-llm?tab=readme-ov-file#install-ipex-llm."
- )
-
- _model_kwargs = model_kwargs or {}
- kwargs = kwargs or {}
-
- _tokenizer_id = tokenizer_id or model_id
- # Set "cpu" as default device
- if "device" not in _model_kwargs:
- _model_kwargs["device"] = "cpu"
-
- if _model_kwargs["device"] not in ["cpu", "xpu"]:
- raise ValueError(
- "IpexLLMBgeEmbeddings currently only supports device to be "
- f"'cpu' or 'xpu', but you have: {_model_kwargs['device']}."
- )
- device = _model_kwargs.pop("device")
-
- try:
- tokenizer = AutoTokenizer.from_pretrained(_tokenizer_id, **_model_kwargs)
- except Exception:
- tokenizer = LlamaTokenizer.from_pretrained(_tokenizer_id, **_model_kwargs)
-
- # restore model_kwargs
- if "trust_remote_code" in _model_kwargs:
- _model_kwargs = {
- k: v for k, v in _model_kwargs.items() if k != "trust_remote_code"
- }
-
- # load model with AutoModelForCausalLM and falls back to AutoModel on failure.
- load_kwargs = {
- "use_cache": True,
- "trust_remote_code": True,
- }
-
- if not low_bit_model:
- if load_in_low_bit is not None:
- load_function_name = "from_pretrained"
- load_kwargs["load_in_low_bit"] = load_in_low_bit # type: ignore[assignment]
- else:
- load_function_name = "from_pretrained"
- load_kwargs["load_in_4bit"] = load_in_4bit
- else:
- load_function_name = "load_low_bit"
-
- try:
- # Attempt to load with AutoModelForCausalLM
- model = cls._load_model_general(
- AutoModelForCausalLM,
- load_function_name=load_function_name,
- model_id=model_id,
- load_kwargs=load_kwargs,
- model_kwargs=_model_kwargs,
- )
- except Exception:
- # Fallback to AutoModel if there's an exception
- model = cls._load_model_general(
- AutoModel,
- load_function_name=load_function_name,
- model_id=model_id,
- load_kwargs=load_kwargs,
- model_kwargs=_model_kwargs,
- )
-
- model.to(device)
-
- return cls(
- model_id=model_id,
- model=model,
- tokenizer=tokenizer,
- model_kwargs=_model_kwargs,
- **kwargs,
- )
-
- @staticmethod
- def _load_model_general(
- model_class: Any,
- load_function_name: str,
- model_id: str,
- load_kwargs: dict,
- model_kwargs: dict,
- ) -> Any:
- """General function to attempt to load a model."""
- try:
- load_function = getattr(model_class, load_function_name)
- return load_function(model_id, **{**load_kwargs, **model_kwargs})
- except Exception as e:
- logger.error(
- f"Failed to load model using "
- f"{model_class.__name__}.{load_function_name}: {e}"
- )
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- "model_id": self.model_id,
- "model_kwargs": self.model_kwargs,
- }
-
- @property
- def _llm_type(self) -> str:
- return "ipex-llm"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- if self.streaming:
- from transformers import TextStreamer
-
- input_ids = self.tokenizer.encode(prompt, return_tensors="pt")
- input_ids = input_ids.to(self.model.device)
- streamer = TextStreamer(
- self.tokenizer, skip_prompt=True, skip_special_tokens=True
- )
- if stop is not None:
- from transformers.generation.stopping_criteria import (
- StoppingCriteriaList,
- )
- from transformers.tools.agents import StopSequenceCriteria
-
- # stop generation when stop words are encountered
- # TODO: stop generation when the following one is stop word
- stopping_criteria = StoppingCriteriaList(
- [StopSequenceCriteria(stop, self.tokenizer)]
- )
- else:
- stopping_criteria = None
- output = self.model.generate(
- input_ids,
- streamer=streamer,
- stopping_criteria=stopping_criteria,
- **kwargs,
- )
- text = self.tokenizer.decode(output[0], skip_special_tokens=True)
- return text
- else:
- input_ids = self.tokenizer.encode(prompt, return_tensors="pt")
- input_ids = input_ids.to(self.model.device)
- if stop is not None:
- from transformers.generation.stopping_criteria import (
- StoppingCriteriaList,
- )
- from transformers.tools.agents import StopSequenceCriteria
-
- stopping_criteria = StoppingCriteriaList(
- [StopSequenceCriteria(stop, self.tokenizer)]
- )
- else:
- stopping_criteria = None
- output = self.model.generate(
- input_ids, stopping_criteria=stopping_criteria, **kwargs
- )
- text = self.tokenizer.decode(output[0], skip_special_tokens=True)[
- len(prompt) :
- ]
- return text
diff --git a/libs/community/langchain_community/llms/javelin_ai_gateway.py b/libs/community/langchain_community/llms/javelin_ai_gateway.py
deleted file mode 100644
index 3571ea4d8b..0000000000
--- a/libs/community/langchain_community/llms/javelin_ai_gateway.py
+++ /dev/null
@@ -1,151 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Dict, List, Mapping, Optional
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from pydantic import BaseModel
-
-
-# Ignoring type because below is valid pydantic code
-# Unexpected keyword argument "extra" for "__init_subclass__" of "object"
-class Params(BaseModel, extra="allow"):
- """Parameters for the Javelin AI Gateway LLM."""
-
- temperature: float = 0.0
- stop: Optional[List[str]] = None
- max_tokens: Optional[int] = None
-
-
-class JavelinAIGateway(LLM):
- """Javelin AI Gateway LLMs.
-
- To use, you should have the ``javelin_sdk`` python package installed.
- For more information, see https://docs.getjavelin.io
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import JavelinAIGateway
-
- completions = JavelinAIGateway(
- gateway_uri="",
- route="",
- params={
- "temperature": 0.1
- }
- )
- """
-
- route: str
- """The route to use for the Javelin AI Gateway API."""
-
- client: Optional[Any] = None
- """The Javelin AI Gateway client."""
-
- gateway_uri: Optional[str] = None
- """The URI of the Javelin AI Gateway API."""
-
- params: Optional[Params] = None
- """Parameters for the Javelin AI Gateway API."""
-
- javelin_api_key: Optional[str] = None
- """The API key for the Javelin AI Gateway API."""
-
- def __init__(self, **kwargs: Any):
- try:
- from javelin_sdk import (
- JavelinClient,
- UnauthorizedError,
- )
- except ImportError:
- raise ImportError(
- "Could not import javelin_sdk python package. "
- "Please install it with `pip install javelin_sdk`."
- )
- super().__init__(**kwargs)
- if self.gateway_uri:
- try:
- self.client = JavelinClient(
- base_url=self.gateway_uri, api_key=self.javelin_api_key
- )
- except UnauthorizedError as e:
- raise ValueError("Javelin: Incorrect API Key.") from e
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling Javelin AI Gateway API."""
- params: Dict[str, Any] = {
- "gateway_uri": self.gateway_uri,
- "route": self.route,
- "javelin_api_key": self.javelin_api_key,
- **(self.params.dict() if self.params else {}),
- }
- return params
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return self._default_params
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call the Javelin AI Gateway API."""
- data: Dict[str, Any] = {
- "prompt": prompt,
- **(self.params.dict() if self.params else {}),
- }
- if s := (stop or (self.params.stop if self.params else None)):
- data["stop"] = s
-
- if self.client is not None:
- resp = self.client.query_route(self.route, query_body=data)
- else:
- raise ValueError("Javelin client is not initialized.")
-
- resp_dict = resp.dict()
-
- try:
- return resp_dict["llm_response"]["choices"][0]["text"]
- except KeyError:
- return ""
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call async the Javelin AI Gateway API."""
- data: Dict[str, Any] = {
- "prompt": prompt,
- **(self.params.dict() if self.params else {}),
- }
- if s := (stop or (self.params.stop if self.params else None)):
- data["stop"] = s
-
- if self.client is not None:
- resp = await self.client.aquery_route(self.route, query_body=data)
- else:
- raise ValueError("Javelin client is not initialized.")
-
- resp_dict = resp.dict()
-
- try:
- return resp_dict["llm_response"]["choices"][0]["text"]
- except KeyError:
- return ""
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "javelin-ai-gateway"
diff --git a/libs/community/langchain_community/llms/koboldai.py b/libs/community/langchain_community/llms/koboldai.py
deleted file mode 100644
index 837dd306ca..0000000000
--- a/libs/community/langchain_community/llms/koboldai.py
+++ /dev/null
@@ -1,197 +0,0 @@
-import logging
-from typing import Any, Dict, List, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-
-logger = logging.getLogger(__name__)
-
-
-def clean_url(url: str) -> str:
- """Remove trailing slash and /api from url if present."""
- if url.endswith("/api"):
- return url[:-4]
- elif url.endswith("/"):
- return url[:-1]
- else:
- return url
-
-
-class KoboldApiLLM(LLM):
- """Kobold API language model.
-
- It includes several fields that can be used to control the text generation process.
-
- To use this class, instantiate it with the required parameters and call it with a
- prompt to generate text. For example:
-
- kobold = KoboldApiLLM(endpoint="http://localhost:5000")
- result = kobold("Write a story about a dragon.")
-
- This will send a POST request to the Kobold API with the provided prompt and
- generate text.
- """
-
- endpoint: str
- """The API endpoint to use for generating text."""
-
- use_story: Optional[bool] = False
- """ Whether or not to use the story from the KoboldAI GUI when generating text. """
-
- use_authors_note: Optional[bool] = False
- """Whether to use the author's note from the KoboldAI GUI when generating text.
-
- This has no effect unless use_story is also enabled.
- """
-
- use_world_info: Optional[bool] = False
- """Whether to use the world info from the KoboldAI GUI when generating text."""
-
- use_memory: Optional[bool] = False
- """Whether to use the memory from the KoboldAI GUI when generating text."""
-
- max_context_length: Optional[int] = 1600
- """Maximum number of tokens to send to the model.
-
- minimum: 1
- """
-
- max_length: Optional[int] = 80
- """Number of tokens to generate.
-
- maximum: 512
- minimum: 1
- """
-
- rep_pen: Optional[float] = 1.12
- """Base repetition penalty value.
-
- minimum: 1
- """
-
- rep_pen_range: Optional[int] = 1024
- """Repetition penalty range.
-
- minimum: 0
- """
-
- rep_pen_slope: Optional[float] = 0.9
- """Repetition penalty slope.
-
- minimum: 0
- """
-
- temperature: Optional[float] = 0.6
- """Temperature value.
-
- exclusiveMinimum: 0
- """
-
- tfs: Optional[float] = 0.9
- """Tail free sampling value.
-
- maximum: 1
- minimum: 0
- """
-
- top_a: Optional[float] = 0.9
- """Top-a sampling value.
-
- minimum: 0
- """
-
- top_p: Optional[float] = 0.95
- """Top-p sampling value.
-
- maximum: 1
- minimum: 0
- """
-
- top_k: Optional[int] = 0
- """Top-k sampling value.
-
- minimum: 0
- """
-
- typical: Optional[float] = 0.5
- """Typical sampling value.
-
- maximum: 1
- minimum: 0
- """
-
- @property
- def _llm_type(self) -> str:
- return "koboldai"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call the API and return the output.
-
- Args:
- prompt: The prompt to use for generation.
- stop: A list of strings to stop generation when encountered.
-
- Returns:
- The generated text.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import KoboldApiLLM
-
- llm = KoboldApiLLM(endpoint="http://localhost:5000")
- llm.invoke("Write a story about dragons.")
- """
- data: Dict[str, Any] = {
- "prompt": prompt,
- "use_story": self.use_story,
- "use_authors_note": self.use_authors_note,
- "use_world_info": self.use_world_info,
- "use_memory": self.use_memory,
- "max_context_length": self.max_context_length,
- "max_length": self.max_length,
- "rep_pen": self.rep_pen,
- "rep_pen_range": self.rep_pen_range,
- "rep_pen_slope": self.rep_pen_slope,
- "temperature": self.temperature,
- "tfs": self.tfs,
- "top_a": self.top_a,
- "top_p": self.top_p,
- "top_k": self.top_k,
- "typical": self.typical,
- }
-
- if stop is not None:
- data["stop_sequence"] = stop
-
- response = requests.post(
- f"{clean_url(self.endpoint)}/api/v1/generate", json=data
- )
-
- response.raise_for_status()
- json_response = response.json()
-
- if (
- "results" in json_response
- and len(json_response["results"]) > 0
- and "text" in json_response["results"][0]
- ):
- text = json_response["results"][0]["text"].strip()
-
- if stop is not None:
- for sequence in stop:
- if text.endswith(sequence):
- text = text[: -len(sequence)].rstrip()
-
- return text
- else:
- raise ValueError(
- f"Unexpected response format from Kobold API: {json_response}"
- )
diff --git a/libs/community/langchain_community/llms/konko.py b/libs/community/langchain_community/llms/konko.py
deleted file mode 100644
index 0c2a62e927..0000000000
--- a/libs/community/langchain_community/llms/konko.py
+++ /dev/null
@@ -1,201 +0,0 @@
-"""Wrapper around Konko AI's Completion API."""
-
-import logging
-import warnings
-from typing import Any, Dict, List, Optional
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from pydantic import ConfigDict, SecretStr, model_validator
-
-from langchain_community.utils.openai import is_openai_v1
-
-logger = logging.getLogger(__name__)
-
-
-class Konko(LLM):
- """Konko AI models.
-
- To use, you'll need an API key. This can be passed in as init param
- ``konko_api_key`` or set as environment variable ``KONKO_API_KEY``.
-
- Konko AI API reference: https://docs.konko.ai/reference/
- """
-
- base_url: str = "https://api.konko.ai/v1/completions"
- """Base inference API URL."""
- konko_api_key: SecretStr
- """Konko AI API key."""
- model: str
- """Model name. Available models listed here:
- https://docs.konko.ai/reference/get_models
- """
- temperature: Optional[float] = None
- """Model temperature."""
- top_p: Optional[float] = None
- """Used to dynamically adjust the number of choices for each predicted token based
- on the cumulative probabilities. A value of 1 will always yield the same
- output. A temperature less than 1 favors more correctness and is appropriate
- for question answering or summarization. A value greater than 1 introduces more
- randomness in the output.
- """
- top_k: Optional[int] = None
- """Used to limit the number of choices for the next predicted word or token. It
- specifies the maximum number of tokens to consider at each step, based on their
- probability of occurrence. This technique helps to speed up the generation
- process and can improve the quality of the generated text by focusing on the
- most likely options.
- """
- max_tokens: Optional[int] = None
- """The maximum number of tokens to generate."""
- repetition_penalty: Optional[float] = None
- """A number that controls the diversity of generated text by reducing the
- likelihood of repeated sequences. Higher values decrease repetition.
- """
- logprobs: Optional[int] = None
- """An integer that specifies how many top token log probabilities are included in
- the response for each token generation step.
- """
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict[str, Any]) -> Any:
- """Validate that python package exists in environment."""
- try:
- import konko
-
- except ImportError:
- raise ImportError(
- "Could not import konko python package. "
- "Please install it with `pip install konko`."
- )
- if not hasattr(konko, "_is_legacy_openai"):
- warnings.warn(
- "You are using an older version of the 'konko' package. "
- "Please consider upgrading to access new features"
- "including the completion endpoint."
- )
- return values
-
- def construct_payload(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- **kwargs: Any,
- ) -> Dict[str, Any]:
- stop_to_use = stop[0] if stop and len(stop) == 1 else stop
- payload: Dict[str, Any] = {
- **self.default_params,
- "prompt": prompt,
- "stop": stop_to_use,
- **kwargs,
- }
- return {k: v for k, v in payload.items() if v is not None}
-
- @property
- def _llm_type(self) -> str:
- """Return type of model."""
- return "konko"
-
- @staticmethod
- def get_user_agent() -> str:
- from langchain_community import __version__
-
- return f"langchain/{__version__}"
-
- @property
- def default_params(self) -> Dict[str, Any]:
- return {
- "model": self.model,
- "temperature": self.temperature,
- "top_p": self.top_p,
- "top_k": self.top_k,
- "max_tokens": self.max_tokens,
- "repetition_penalty": self.repetition_penalty,
- }
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Konko's text generation endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
-
- Returns:
- The string generated by the model..
- """
- import konko
-
- payload = self.construct_payload(prompt, stop, **kwargs)
-
- try:
- if is_openai_v1():
- response = konko.completions.create(**payload)
- else:
- response = konko.Completion.create(**payload)
-
- except AttributeError:
- raise ValueError(
- "`konko` has no `Completion` attribute, this is likely "
- "due to an old version of the konko package. Try upgrading it "
- "with `pip install --upgrade konko`."
- )
-
- if is_openai_v1():
- output = response.choices[0].text
- else:
- output = response["choices"][0]["text"]
-
- return output
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Asynchronously call out to Konko's text generation endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
-
- Returns:
- The string generated by the model.
- """
- import konko
-
- payload = self.construct_payload(prompt, stop, **kwargs)
-
- try:
- if is_openai_v1():
- client = konko.AsyncKonko()
- response = await client.completions.create(**payload)
- else:
- response = await konko.Completion.acreate(**payload)
-
- except AttributeError:
- raise ValueError(
- "`konko` has no `Completion` attribute, this is likely "
- "due to an old version of the konko package. Try upgrading it "
- "with `pip install --upgrade konko`."
- )
-
- if is_openai_v1():
- output = response.choices[0].text
- else:
- output = response["choices"][0]["text"]
-
- return output
diff --git a/libs/community/langchain_community/llms/layerup_security.py b/libs/community/langchain_community/llms/layerup_security.py
deleted file mode 100644
index c15626eff2..0000000000
--- a/libs/community/langchain_community/llms/layerup_security.py
+++ /dev/null
@@ -1,107 +0,0 @@
-import logging
-from typing import Any, Callable, Dict, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from pydantic import model_validator
-
-logger = logging.getLogger(__name__)
-
-
-def default_guardrail_violation_handler(violation: dict) -> str:
- """Default guardrail violation handler.
-
- Args:
- violation (dict): The violation dictionary.
-
- Returns:
- str: The canned response.
- """
- if violation.get("canned_response"):
- return violation["canned_response"]
- guardrail_name = (
- f"Guardrail {violation.get('offending_guardrail')}"
- if violation.get("offending_guardrail")
- else "A guardrail"
- )
- raise ValueError(
- f"{guardrail_name} was violated without a proper guardrail violation handler."
- )
-
-
-class LayerupSecurity(LLM):
- """Layerup Security LLM service."""
-
- llm: LLM
- layerup_api_key: str
- layerup_api_base_url: str = "https://api.uselayerup.com/v1"
- prompt_guardrails: Optional[List[str]] = []
- response_guardrails: Optional[List[str]] = []
- mask: bool = False
- metadata: Optional[Dict[str, Any]] = {}
- handle_prompt_guardrail_violation: Callable[[dict], str] = (
- default_guardrail_violation_handler
- )
- handle_response_guardrail_violation: Callable[[dict], str] = (
- default_guardrail_violation_handler
- )
- client: Any #: :meta private:
-
- @model_validator(mode="before")
- @classmethod
- def validate_layerup_sdk(cls, values: Dict[str, Any]) -> Any:
- try:
- from layerup_security import LayerupSecurity as LayerupSecuritySDK
-
- values["client"] = LayerupSecuritySDK(
- api_key=values["layerup_api_key"],
- base_url=values["layerup_api_base_url"],
- )
- except ImportError:
- raise ImportError(
- "Could not import LayerupSecurity SDK. "
- "Please install it with `pip install LayerupSecurity`."
- )
- return values
-
- @property
- def _llm_type(self) -> str:
- return "layerup_security"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- messages = [{"role": "user", "content": prompt}]
- unmask_response = None
-
- if self.mask:
- messages, unmask_response = self.client.mask_prompt(messages, self.metadata)
-
- if self.prompt_guardrails:
- security_response = self.client.execute_guardrails(
- self.prompt_guardrails, messages, prompt, self.metadata
- )
- if not security_response["all_safe"]:
- return self.handle_prompt_guardrail_violation(security_response)
-
- result = self.llm._call(
- messages[0]["content"], run_manager=run_manager, **kwargs
- )
-
- if self.mask and unmask_response:
- result = unmask_response(result)
-
- messages.append({"role": "assistant", "content": result})
-
- if self.response_guardrails:
- security_response = self.client.execute_guardrails(
- self.response_guardrails, messages, result, self.metadata
- )
- if not security_response["all_safe"]:
- return self.handle_response_guardrail_violation(security_response)
-
- return result
diff --git a/libs/community/langchain_community/llms/llamacpp.py b/libs/community/langchain_community/llms/llamacpp.py
deleted file mode 100644
index a045878fd2..0000000000
--- a/libs/community/langchain_community/llms/llamacpp.py
+++ /dev/null
@@ -1,353 +0,0 @@
-from __future__ import annotations
-
-import logging
-from pathlib import Path
-from typing import Any, Dict, Iterator, List, Optional, Union
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import get_pydantic_field_names, pre_init
-from langchain_core.utils.utils import _build_model_kwargs
-from pydantic import Field, model_validator
-
-logger = logging.getLogger(__name__)
-
-
-class LlamaCpp(LLM):
- """llama.cpp model.
-
- To use, you should have the llama-cpp-python library installed, and provide the
- path to the Llama model as a named parameter to the constructor.
- Check out: https://github.com/abetlen/llama-cpp-python
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import LlamaCpp
- llm = LlamaCpp(model_path="/path/to/llama/model")
- """
-
- client: Any = None #: :meta private:
- model_path: str
- """The path to the Llama model file."""
-
- lora_base: Optional[str] = None
- """The path to the Llama LoRA base model."""
-
- lora_path: Optional[str] = None
- """The path to the Llama LoRA. If None, no LoRa is loaded."""
-
- n_ctx: int = Field(512, alias="n_ctx")
- """Token context window."""
-
- n_parts: int = Field(-1, alias="n_parts")
- """Number of parts to split the model into.
- If -1, the number of parts is automatically determined."""
-
- seed: int = Field(-1, alias="seed")
- """Seed. If -1, a random seed is used."""
-
- f16_kv: bool = Field(True, alias="f16_kv")
- """Use half-precision for key/value cache."""
-
- logits_all: bool = Field(False, alias="logits_all")
- """Return logits for all tokens, not just the last token."""
-
- vocab_only: bool = Field(False, alias="vocab_only")
- """Only load the vocabulary, no weights."""
-
- use_mlock: bool = Field(False, alias="use_mlock")
- """Force system to keep model in RAM."""
-
- n_threads: Optional[int] = Field(None, alias="n_threads")
- """Number of threads to use.
- If None, the number of threads is automatically determined."""
-
- n_batch: Optional[int] = Field(8, alias="n_batch")
- """Number of tokens to process in parallel.
- Should be a number between 1 and n_ctx."""
-
- n_gpu_layers: Optional[int] = Field(None, alias="n_gpu_layers")
- """Number of layers to be loaded into gpu memory. Default None."""
-
- suffix: Optional[str] = Field(None)
- """A suffix to append to the generated text. If None, no suffix is appended."""
-
- max_tokens: Optional[int] = 256
- """The maximum number of tokens to generate."""
-
- temperature: Optional[float] = 0.8
- """The temperature to use for sampling."""
-
- top_p: Optional[float] = 0.95
- """The top-p value to use for sampling."""
-
- logprobs: Optional[int] = Field(None)
- """The number of logprobs to return. If None, no logprobs are returned."""
-
- echo: Optional[bool] = False
- """Whether to echo the prompt."""
-
- stop: Optional[List[str]] = []
- """A list of strings to stop generation when encountered."""
-
- repeat_penalty: Optional[float] = 1.1
- """The penalty to apply to repeated tokens."""
-
- top_k: Optional[int] = 40
- """The top-k value to use for sampling."""
-
- last_n_tokens_size: Optional[int] = 64
- """The number of tokens to look back when applying the repeat_penalty."""
-
- use_mmap: Optional[bool] = True
- """Whether to keep the model loaded in RAM"""
-
- rope_freq_scale: float = 1.0
- """Scale factor for rope sampling."""
-
- rope_freq_base: float = 10000.0
- """Base frequency for rope sampling."""
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Any additional parameters to pass to llama_cpp.Llama."""
-
- streaming: bool = True
- """Whether to stream the results, token by token."""
-
- grammar_path: Optional[Union[str, Path]] = None
- """
- grammar_path: Path to the .gbnf file that defines formal grammars
- for constraining model outputs. For instance, the grammar can be used
- to force the model to generate valid JSON or to speak exclusively in emojis. At most
- one of grammar_path and grammar should be passed in.
- """
- grammar: Optional[Union[str, Any]] = None
- """
- grammar: formal grammar for constraining model outputs. For instance, the grammar
- can be used to force the model to generate valid JSON or to speak exclusively in
- emojis. At most one of grammar_path and grammar should be passed in.
- """
-
- verbose: bool = True
- """Print verbose output to stderr."""
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that llama-cpp-python library is installed."""
- try:
- from llama_cpp import Llama, LlamaGrammar
- except ImportError:
- raise ImportError(
- "Could not import llama-cpp-python library. "
- "Please install the llama-cpp-python library to "
- "use this embedding model: pip install llama-cpp-python"
- )
-
- model_path = values["model_path"]
- model_param_names = [
- "rope_freq_scale",
- "rope_freq_base",
- "lora_path",
- "lora_base",
- "n_ctx",
- "n_parts",
- "seed",
- "f16_kv",
- "logits_all",
- "vocab_only",
- "use_mlock",
- "n_threads",
- "n_batch",
- "use_mmap",
- "last_n_tokens_size",
- "verbose",
- ]
- model_params = {k: values[k] for k in model_param_names}
- # For backwards compatibility, only include if non-null.
- if values["n_gpu_layers"] is not None:
- model_params["n_gpu_layers"] = values["n_gpu_layers"]
-
- model_params.update(values["model_kwargs"])
-
- try:
- values["client"] = Llama(model_path, **model_params)
- except Exception as e:
- raise ValueError(
- f"Could not load Llama model from path: {model_path}. "
- f"Received error {e}"
- )
-
- if values["grammar"] and values["grammar_path"]:
- grammar = values["grammar"]
- grammar_path = values["grammar_path"]
- raise ValueError(
- "Can only pass in one of grammar and grammar_path. Received "
- f"{grammar=} and {grammar_path=}."
- )
- elif isinstance(values["grammar"], str):
- values["grammar"] = LlamaGrammar.from_string(values["grammar"])
- elif values["grammar_path"]:
- values["grammar"] = LlamaGrammar.from_file(values["grammar_path"])
- else:
- pass
- return values
-
- @model_validator(mode="before")
- @classmethod
- def build_model_kwargs(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = get_pydantic_field_names(cls)
- values = _build_model_kwargs(values, all_required_field_names)
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling llama_cpp."""
- params = {
- "suffix": self.suffix,
- "max_tokens": self.max_tokens,
- "temperature": self.temperature,
- "top_p": self.top_p,
- "logprobs": self.logprobs,
- "echo": self.echo,
- "stop_sequences": self.stop, # key here is convention among LLM classes
- "repeat_penalty": self.repeat_penalty,
- "top_k": self.top_k,
- }
- if self.grammar:
- params["grammar"] = self.grammar
- return params
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {**{"model_path": self.model_path}, **self._default_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "llamacpp"
-
- def _get_parameters(self, stop: Optional[List[str]] = None) -> Dict[str, Any]:
- """
- Performs sanity check, preparing parameters in format needed by llama_cpp.
-
- Args:
- stop (Optional[List[str]]): List of stop sequences for llama_cpp.
-
- Returns:
- Dictionary containing the combined parameters.
- """
-
- # Raise error if stop sequences are in both input and default params
- if self.stop and stop is not None:
- raise ValueError("`stop` found in both the input and default params.")
-
- params = self._default_params
-
- # llama_cpp expects the "stop" key not this, so we remove it:
- params.pop("stop_sequences")
-
- # then sets it as configured, or default to an empty list:
- params["stop"] = self.stop or stop or []
-
- return params
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call the Llama model and return the output.
-
- Args:
- prompt: The prompt to use for generation.
- stop: A list of strings to stop generation when encountered.
-
- Returns:
- The generated text.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import LlamaCpp
- llm = LlamaCpp(model_path="/path/to/local/llama/model.bin")
- llm.invoke("This is a prompt.")
- """
- if self.streaming:
- # If streaming is enabled, we use the stream
- # method that yields as they are generated
- # and return the combined strings from the first choices's text:
- combined_text_output = ""
- for chunk in self._stream(
- prompt=prompt,
- stop=stop,
- run_manager=run_manager,
- **kwargs,
- ):
- combined_text_output += chunk.text
- return combined_text_output
- else:
- params = self._get_parameters(stop)
- params = {**params, **kwargs}
- result = self.client(prompt=prompt, **params)
- return result["choices"][0]["text"]
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- """Yields results objects as they are generated in real time.
-
- It also calls the callback manager's on_llm_new_token event with
- similar parameters to the OpenAI LLM class method of the same name.
-
- Args:
- prompt: The prompts to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- A generator representing the stream of tokens being generated.
-
- Yields:
- A dictionary like objects containing a string token and metadata.
- See llama-cpp-python docs and below for more.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import LlamaCpp
- llm = LlamaCpp(
- model_path="/path/to/local/model.bin",
- temperature = 0.5
- )
- for chunk in llm.stream("Ask 'Hi, how are you?' like a pirate:'",
- stop=["'","\n"]):
- result = chunk["choices"][0]
- print(result["text"], end='', flush=True) # noqa: T201
-
- """
- params = {**self._get_parameters(stop), **kwargs}
- result = self.client(prompt=prompt, stream=True, **params)
- for part in result:
- logprobs = part["choices"][0].get("logprobs", None)
- chunk = GenerationChunk(
- text=part["choices"][0]["text"],
- generation_info={"logprobs": logprobs},
- )
- if run_manager:
- run_manager.on_llm_new_token(
- token=chunk.text, verbose=self.verbose, log_probs=logprobs
- )
- yield chunk
-
- def get_num_tokens(self, text: str) -> int:
- tokenized_text = self.client.tokenize(text.encode("utf-8"))
- return len(tokenized_text)
diff --git a/libs/community/langchain_community/llms/llamafile.py b/libs/community/langchain_community/llms/llamafile.py
deleted file mode 100644
index e168572c1b..0000000000
--- a/libs/community/langchain_community/llms/llamafile.py
+++ /dev/null
@@ -1,319 +0,0 @@
-from __future__ import annotations
-
-import json
-from io import StringIO
-from typing import Any, Dict, Iterator, List, Optional
-
-import requests
-from langchain_core.callbacks.manager import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import get_pydantic_field_names
-from pydantic import ConfigDict
-
-
-class Llamafile(LLM):
- """Llamafile lets you distribute and run large language models with a
- single file.
-
- To get started, see: https://github.com/Mozilla-Ocho/llamafile
-
- To use this class, you will need to first:
-
- 1. Download a llamafile.
- 2. Make the downloaded file executable: `chmod +x path/to/model.llamafile`
- 3. Start the llamafile in server mode:
-
- `./path/to/model.llamafile --server --nobrowser`
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import Llamafile
- llm = Llamafile()
- llm.invoke("Tell me a joke.")
- """
-
- base_url: str = "http://localhost:8080"
- """Base url where the llamafile server is listening."""
-
- request_timeout: Optional[int] = None
- """Timeout for server requests"""
-
- streaming: bool = False
- """Allows receiving each predicted token in real-time instead of
- waiting for the completion to finish. To enable this, set to true."""
-
- # Generation options
-
- seed: int = -1
- """Random Number Generator (RNG) seed. A random seed is used if this is
- less than zero. Default: -1"""
-
- temperature: float = 0.8
- """Temperature. Default: 0.8"""
-
- top_k: int = 40
- """Limit the next token selection to the K most probable tokens.
- Default: 40."""
-
- top_p: float = 0.95
- """Limit the next token selection to a subset of tokens with a cumulative
- probability above a threshold P. Default: 0.95."""
-
- min_p: float = 0.05
- """The minimum probability for a token to be considered, relative to
- the probability of the most likely token. Default: 0.05."""
-
- n_predict: int = -1
- """Set the maximum number of tokens to predict when generating text.
- Note: May exceed the set limit slightly if the last token is a partial
- multibyte character. When 0, no tokens will be generated but the prompt
- is evaluated into the cache. Default: -1 = infinity."""
-
- n_keep: int = 0
- """Specify the number of tokens from the prompt to retain when the
- context size is exceeded and tokens need to be discarded. By default,
- this value is set to 0 (meaning no tokens are kept). Use -1 to retain all
- tokens from the prompt."""
-
- tfs_z: float = 1.0
- """Enable tail free sampling with parameter z. Default: 1.0 = disabled."""
-
- typical_p: float = 1.0
- """Enable locally typical sampling with parameter p.
- Default: 1.0 = disabled."""
-
- repeat_penalty: float = 1.1
- """Control the repetition of token sequences in the generated text.
- Default: 1.1"""
-
- repeat_last_n: int = 64
- """Last n tokens to consider for penalizing repetition. Default: 64,
- 0 = disabled, -1 = ctx-size."""
-
- penalize_nl: bool = True
- """Penalize newline tokens when applying the repeat penalty.
- Default: true."""
-
- presence_penalty: float = 0.0
- """Repeat alpha presence penalty. Default: 0.0 = disabled."""
-
- frequency_penalty: float = 0.0
- """Repeat alpha frequency penalty. Default: 0.0 = disabled"""
-
- mirostat: int = 0
- """Enable Mirostat sampling, controlling perplexity during text
- generation. 0 = disabled, 1 = Mirostat, 2 = Mirostat 2.0.
- Default: disabled."""
-
- mirostat_tau: float = 5.0
- """Set the Mirostat target entropy, parameter tau. Default: 5.0."""
-
- mirostat_eta: float = 0.1
- """Set the Mirostat learning rate, parameter eta. Default: 0.1."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @property
- def _llm_type(self) -> str:
- return "llamafile"
-
- @property
- def _param_fieldnames(self) -> List[str]:
- # Return the list of fieldnames that will be passed as configurable
- # generation options to the llamafile server. Exclude 'builtin' fields
- # from the BaseLLM class like 'metadata' as well as fields that should
- # not be passed in requests (base_url, request_timeout).
- ignore_keys = [
- "base_url",
- "cache",
- "callback_manager",
- "callbacks",
- "metadata",
- "name",
- "request_timeout",
- "streaming",
- "tags",
- "verbose",
- "custom_get_token_ids",
- ]
- attrs = [
- k for k in get_pydantic_field_names(self.__class__) if k not in ignore_keys
- ]
- return attrs
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- params = {}
- for fieldname in self._param_fieldnames:
- params[fieldname] = getattr(self, fieldname)
- return params
-
- def _get_parameters(
- self, stop: Optional[List[str]] = None, **kwargs: Any
- ) -> Dict[str, Any]:
- params = self._default_params
-
- # Only update keys that are already present in the default config.
- # This way, we don't accidentally post unknown/unhandled key/values
- # in the request to the llamafile server
- for k, v in kwargs.items():
- if k in params:
- params[k] = v
-
- if stop is not None and len(stop) > 0:
- params["stop"] = stop
-
- if self.streaming:
- params["stream"] = True
-
- return params
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Request prompt completion from the llamafile server and return the
- output.
-
- Args:
- prompt: The prompt to use for generation.
- stop: A list of strings to stop generation when encountered.
- run_manager:
- **kwargs: Any additional options to pass as part of the
- generation request.
-
- Returns:
- The string generated by the model.
-
- """
-
- if self.streaming:
- with StringIO() as buff:
- for chunk in self._stream(
- prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- buff.write(chunk.text)
-
- text = buff.getvalue()
-
- return text
-
- else:
- params = self._get_parameters(stop=stop, **kwargs)
- payload = {"prompt": prompt, **params}
-
- try:
- response = requests.post(
- url=f"{self.base_url}/completion",
- headers={
- "Content-Type": "application/json",
- },
- json=payload,
- stream=False,
- timeout=self.request_timeout,
- )
- except requests.exceptions.ConnectionError:
- raise requests.exceptions.ConnectionError(
- f"Could not connect to Llamafile server. Please make sure "
- f"that a server is running at {self.base_url}."
- )
-
- response.raise_for_status()
- response.encoding = "utf-8"
-
- text = response.json()["content"]
-
- return text
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- """Yields results objects as they are generated in real time.
-
- It also calls the callback manager's on_llm_new_token event with
- similar parameters to the OpenAI LLM class method of the same name.
-
- Args:
- prompt: The prompts to pass into the model.
- stop: Optional list of stop words to use when generating.
- run_manager:
- **kwargs: Any additional options to pass as part of the
- generation request.
-
- Returns:
- A generator representing the stream of tokens being generated.
-
- Yields:
- Dictionary-like objects each containing a token
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import Llamafile
- llm = Llamafile(
- temperature = 0.0
- )
- for chunk in llm.stream("Ask 'Hi, how are you?' like a pirate:'",
- stop=["'","\n"]):
- result = chunk["choices"][0]
- print(result["text"], end='', flush=True)
-
- """
- params = self._get_parameters(stop=stop, **kwargs)
- if "stream" not in params:
- params["stream"] = True
-
- payload = {"prompt": prompt, **params}
-
- try:
- response = requests.post(
- url=f"{self.base_url}/completion",
- headers={
- "Content-Type": "application/json",
- },
- json=payload,
- stream=True,
- timeout=self.request_timeout,
- )
- except requests.exceptions.ConnectionError:
- raise requests.exceptions.ConnectionError(
- f"Could not connect to Llamafile server. Please make sure "
- f"that a server is running at {self.base_url}."
- )
-
- response.encoding = "utf8"
-
- for raw_chunk in response.iter_lines(decode_unicode=True):
- content = self._get_chunk_content(raw_chunk)
- chunk = GenerationChunk(text=content)
-
- if run_manager:
- run_manager.on_llm_new_token(token=chunk.text)
- yield chunk
-
- def _get_chunk_content(self, chunk: str) -> str:
- """When streaming is turned on, llamafile server returns lines like:
-
- 'data: {"content":" They","multimodal":true,"slot_id":0,"stop":false}'
-
- Here, we convert this to a dict and return the value of the 'content'
- field
- """
-
- if chunk.startswith("data:"):
- cleaned = chunk.lstrip("data: ")
- data = json.loads(cleaned)
- return data["content"]
- else:
- return chunk
diff --git a/libs/community/langchain_community/llms/loading.py b/libs/community/langchain_community/llms/loading.py
deleted file mode 100644
index 4e97587b1b..0000000000
--- a/libs/community/langchain_community/llms/loading.py
+++ /dev/null
@@ -1,55 +0,0 @@
-"""Base interface for loading large language model APIs."""
-
-import json
-from pathlib import Path
-from typing import Any, Union
-
-import yaml
-from langchain_core.language_models.llms import BaseLLM
-from langchain_core.utils.pydantic import get_fields
-
-from langchain_community.llms import get_type_to_cls_dict
-
-_ALLOW_DANGEROUS_DESERIALIZATION_ARG = "allow_dangerous_deserialization"
-
-
-def load_llm_from_config(config: dict, **kwargs: Any) -> BaseLLM:
- """Load LLM from Config Dict."""
- if "_type" not in config:
- raise ValueError("Must specify an LLM Type in config")
- config_type = config.pop("_type")
-
- type_to_cls_dict = get_type_to_cls_dict()
-
- if config_type not in type_to_cls_dict:
- raise ValueError(f"Loading {config_type} LLM not supported")
-
- llm_cls = type_to_cls_dict[config_type]()
-
- load_kwargs = {}
- if _ALLOW_DANGEROUS_DESERIALIZATION_ARG in get_fields(llm_cls):
- load_kwargs[_ALLOW_DANGEROUS_DESERIALIZATION_ARG] = kwargs.get(
- _ALLOW_DANGEROUS_DESERIALIZATION_ARG, False
- )
-
- return llm_cls(**config, **load_kwargs)
-
-
-def load_llm(file: Union[str, Path], **kwargs: Any) -> BaseLLM:
- """Load LLM from a file."""
- # Convert file to Path object.
- if isinstance(file, str):
- file_path = Path(file)
- else:
- file_path = file
- # Load from either json or yaml.
- if file_path.suffix == ".json":
- with open(file_path) as f:
- config = json.load(f)
- elif file_path.suffix.endswith((".yaml", ".yml")):
- with open(file_path, "r") as f:
- config = yaml.safe_load(f)
- else:
- raise ValueError("File type must be json or yaml")
- # Load the LLM from the config now.
- return load_llm_from_config(config, **kwargs)
diff --git a/libs/community/langchain_community/llms/manifest.py b/libs/community/langchain_community/llms/manifest.py
deleted file mode 100644
index 966933d5f9..0000000000
--- a/libs/community/langchain_community/llms/manifest.py
+++ /dev/null
@@ -1,63 +0,0 @@
-from typing import Any, Dict, List, Mapping, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import pre_init
-from pydantic import ConfigDict
-
-
-class ManifestWrapper(LLM):
- """HazyResearch's Manifest library."""
-
- client: Any = None #: :meta private:
- llm_kwargs: Optional[Dict] = None
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that python package exists in environment."""
- try:
- from manifest import Manifest
-
- if not isinstance(values["client"], Manifest):
- raise ValueError
- except ImportError:
- raise ImportError(
- "Could not import manifest python package. "
- "Please install it with `pip install manifest-ml`."
- )
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- kwargs = self.llm_kwargs or {}
- return {
- **self.client.client_pool.get_current_client().get_model_params(),
- **kwargs,
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "manifest"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to LLM through Manifest."""
- if stop is not None and len(stop) != 1:
- raise NotImplementedError(
- f"Manifest currently only supports a single stop token, got {stop}"
- )
- params = self.llm_kwargs or {}
- params = {**params, **kwargs}
- if stop is not None:
- params["stop_token"] = stop
- return self.client.run(prompt, **params)
diff --git a/libs/community/langchain_community/llms/minimax.py b/libs/community/langchain_community/llms/minimax.py
deleted file mode 100644
index 5a1822fffb..0000000000
--- a/libs/community/langchain_community/llms/minimax.py
+++ /dev/null
@@ -1,161 +0,0 @@
-"""Wrapper around Minimax APIs."""
-
-from __future__ import annotations
-
-import logging
-from typing import (
- Any,
- Dict,
- List,
- Optional,
-)
-
-import requests
-from langchain_core.callbacks import (
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict, Field, SecretStr, model_validator
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-class _MinimaxEndpointClient(BaseModel):
- """API client for the Minimax LLM endpoint."""
-
- host: str
- group_id: str
- api_key: SecretStr
- api_url: str
-
- @model_validator(mode="before")
- @classmethod
- def set_api_url(cls, values: Dict[str, Any]) -> Any:
- if "api_url" not in values:
- host = values["host"]
- group_id = values["group_id"]
- api_url = f"{host}/v1/text/chatcompletion?GroupId={group_id}"
- values["api_url"] = api_url
- return values
-
- def post(self, request: Any) -> Any:
- headers = {"Authorization": f"Bearer {self.api_key.get_secret_value()}"}
- response = requests.post(self.api_url, headers=headers, json=request)
- # TODO: error handling and automatic retries
- if not response.ok:
- raise ValueError(f"HTTP {response.status_code} error: {response.text}")
- if response.json()["base_resp"]["status_code"] > 0:
- raise ValueError(
- f"API {response.json()['base_resp']['status_code']}"
- f" error: {response.json()['base_resp']['status_msg']}"
- )
- return response.json()["reply"]
-
-
-class MinimaxCommon(BaseModel):
- """Common parameters for Minimax large language models."""
-
- model_config = ConfigDict(protected_namespaces=())
-
- _client: _MinimaxEndpointClient
- model: str = "abab5.5-chat"
- """Model name to use."""
- max_tokens: int = 256
- """Denotes the number of tokens to predict per generation."""
- temperature: float = 0.7
- """A non-negative float that tunes the degree of randomness in generation."""
- top_p: float = 0.95
- """Total probability mass of tokens to consider at each step."""
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `create` call not explicitly specified."""
- minimax_api_host: Optional[str] = None
- minimax_group_id: Optional[str] = None
- minimax_api_key: Optional[SecretStr] = None
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- values["minimax_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(values, "minimax_api_key", "MINIMAX_API_KEY")
- )
- values["minimax_group_id"] = get_from_dict_or_env(
- values, "minimax_group_id", "MINIMAX_GROUP_ID"
- )
- # Get custom api url from environment.
- values["minimax_api_host"] = get_from_dict_or_env(
- values,
- "minimax_api_host",
- "MINIMAX_API_HOST",
- default="https://api.minimax.chat",
- )
- values["_client"] = _MinimaxEndpointClient( # type: ignore[call-arg]
- host=values["minimax_api_host"],
- api_key=values["minimax_api_key"],
- group_id=values["minimax_group_id"],
- )
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling OpenAI API."""
- return {
- "model": self.model,
- "tokens_to_generate": self.max_tokens,
- "temperature": self.temperature,
- "top_p": self.top_p,
- **self.model_kwargs,
- }
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {**{"model": self.model}, **self._default_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "minimax"
-
-
-class Minimax(MinimaxCommon, LLM):
- """Minimax large language models.
-
- To use, you should have the environment variable
- ``MINIMAX_API_KEY`` and ``MINIMAX_GROUP_ID`` set with your API key,
- or pass them as a named parameter to the constructor.
- Example:
- . code-block:: python
- from langchain_community.llms.minimax import Minimax
- minimax = Minimax(model="", minimax_api_key="my-api-key",
- minimax_group_id="my-group-id")
- """
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- r"""Call out to Minimax's completion endpoint to chat
- Args:
- prompt: The prompt to pass into the model.
- Returns:
- The string generated by the model.
- Example:
- .. code-block:: python
- response = minimax("Tell me a joke.")
- """
- request = self._default_params
- request["messages"] = [{"sender_type": "USER", "text": prompt}]
- request.update(kwargs)
- text = self._client.post(request)
- if stop is not None:
- # This is required since the stop tokens
- # are not enforced by the model parameters
- text = enforce_stop_tokens(text, stop)
-
- return text
diff --git a/libs/community/langchain_community/llms/mlflow.py b/libs/community/langchain_community/llms/mlflow.py
deleted file mode 100644
index d8422629bc..0000000000
--- a/libs/community/langchain_community/llms/mlflow.py
+++ /dev/null
@@ -1,111 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Dict, List, Mapping, Optional
-from urllib.parse import urlparse
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models import LLM
-from pydantic import Field, PrivateAttr
-
-
-class Mlflow(LLM):
- """MLflow LLM service.
-
- To use, you should have the `mlflow[genai]` python package installed.
- For more information, see https://mlflow.org/docs/latest/llms/deployments.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import Mlflow
-
- completions = Mlflow(
- target_uri="http://localhost:5000",
- endpoint="test",
- temperature=0.1,
- )
- """
-
- endpoint: str
- """The endpoint to use."""
- target_uri: str
- """The target URI to use."""
- temperature: float = 0.0
- """The sampling temperature."""
- n: int = 1
- """The number of completion choices to generate."""
- stop: Optional[List[str]] = None
- """The stop sequence."""
- max_tokens: Optional[int] = None
- """The maximum number of tokens to generate."""
- extra_params: Dict[str, Any] = Field(default_factory=dict)
- """Any extra parameters to pass to the endpoint."""
-
- """Extra parameters such as `temperature`."""
- _client: Any = PrivateAttr()
-
- def __init__(self, **kwargs: Any):
- super().__init__(**kwargs)
- self._validate_uri()
- try:
- from mlflow.deployments import get_deploy_client
-
- self._client = get_deploy_client(self.target_uri)
- except ImportError as e:
- raise ImportError(
- "Failed to create the client. "
- "Please run `pip install mlflow[genai]` to install "
- "required dependencies."
- ) from e
-
- def _validate_uri(self) -> None:
- if self.target_uri == "databricks":
- return
- allowed = ["http", "https", "databricks"]
- if urlparse(self.target_uri).scheme not in allowed:
- raise ValueError(
- f"Invalid target URI: {self.target_uri}. "
- f"The scheme must be one of {allowed}."
- )
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- return {
- "target_uri": self.target_uri,
- "endpoint": self.endpoint,
- "temperature": self.temperature,
- "n": self.n,
- "stop": self.stop,
- "max_tokens": self.max_tokens,
- "extra_params": self.extra_params,
- }
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- return self._default_params
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- data: Dict[str, Any] = {
- "prompt": prompt,
- "temperature": self.temperature,
- "n": self.n,
- **self.extra_params,
- **kwargs,
- }
- if stop := self.stop or stop:
- data["stop"] = stop
- if self.max_tokens is not None:
- data["max_tokens"] = self.max_tokens
-
- resp = self._client.predict(endpoint=self.endpoint, inputs=data)
- return resp["choices"][0]["text"]
-
- @property
- def _llm_type(self) -> str:
- return "mlflow"
diff --git a/libs/community/langchain_community/llms/mlflow_ai_gateway.py b/libs/community/langchain_community/llms/mlflow_ai_gateway.py
deleted file mode 100644
index 8594aab8c5..0000000000
--- a/libs/community/langchain_community/llms/mlflow_ai_gateway.py
+++ /dev/null
@@ -1,103 +0,0 @@
-from __future__ import annotations
-
-import warnings
-from typing import Any, Dict, List, Mapping, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from pydantic import BaseModel
-
-
-# Ignoring type because below is valid pydantic code
-# Unexpected keyword argument "extra" for "__init_subclass__" of "object"
-class Params(BaseModel, extra="allow"):
- """Parameters for the MLflow AI Gateway LLM."""
-
- temperature: float = 0.0
- candidate_count: int = 1
- """The number of candidates to return."""
- stop: Optional[List[str]] = None
- max_tokens: Optional[int] = None
-
-
-class MlflowAIGateway(LLM):
- """MLflow AI Gateway LLMs.
-
- To use, you should have the ``mlflow[gateway]`` python package installed.
- For more information, see https://mlflow.org/docs/latest/gateway/index.html.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import MlflowAIGateway
-
- completions = MlflowAIGateway(
- gateway_uri="",
- route="",
- params={
- "temperature": 0.1
- }
- )
- """
-
- route: str
- gateway_uri: Optional[str] = None
- params: Optional[Params] = None
-
- def __init__(self, **kwargs: Any):
- warnings.warn(
- "`MlflowAIGateway` is deprecated. Use `Mlflow` or `Databricks` instead.",
- DeprecationWarning,
- )
- try:
- import mlflow.gateway
- except ImportError as e:
- raise ImportError(
- "Could not import `mlflow.gateway` module. "
- "Please install it with `pip install mlflow[gateway]`."
- ) from e
-
- super().__init__(**kwargs)
- if self.gateway_uri:
- mlflow.gateway.set_gateway_uri(self.gateway_uri)
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- params: Dict[str, Any] = {
- "gateway_uri": self.gateway_uri,
- "route": self.route,
- **(self.params.dict() if self.params else {}),
- }
- return params
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- return self._default_params
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- try:
- import mlflow.gateway
- except ImportError as e:
- raise ImportError(
- "Could not import `mlflow.gateway` module. "
- "Please install it with `pip install mlflow[gateway]`."
- ) from e
-
- data: Dict[str, Any] = {
- "prompt": prompt,
- **(self.params.dict() if self.params else {}),
- }
- if s := (stop or (self.params.stop if self.params else None)):
- data["stop"] = s
- resp = mlflow.gateway.query(self.route, data=data)
- return resp["candidates"][0]["text"]
-
- @property
- def _llm_type(self) -> str:
- return "mlflow-ai-gateway"
diff --git a/libs/community/langchain_community/llms/mlx_pipeline.py b/libs/community/langchain_community/llms/mlx_pipeline.py
deleted file mode 100644
index 4ef69d4cc0..0000000000
--- a/libs/community/langchain_community/llms/mlx_pipeline.py
+++ /dev/null
@@ -1,256 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Callable, Iterator, List, Mapping, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from pydantic import ConfigDict
-
-DEFAULT_MODEL_ID = "mlx-community/quantized-gemma-2b"
-
-logger = logging.getLogger(__name__)
-
-
-class MLXPipeline(LLM):
- """MLX Pipeline API.
-
- To use, you should have the ``mlx-lm`` python package installed.
-
- Example using from_model_id:
- .. code-block:: python
-
- from langchain_community.llms import MLXPipeline
- pipe = MLXPipeline.from_model_id(
- model_id="mlx-community/quantized-gemma-2b",
- pipeline_kwargs={"max_tokens": 10, "temp": 0.7},
- )
- Example passing model and tokenizer in directly:
- .. code-block:: python
-
- from langchain_community.llms import MLXPipeline
- from mlx_lm import load
- model_id="mlx-community/quantized-gemma-2b"
- model, tokenizer = load(model_id)
- pipe = MLXPipeline(model=model, tokenizer=tokenizer)
- """
-
- model_id: str = DEFAULT_MODEL_ID
- """Model name to use."""
- model: Any = None #: :meta private:
- """Model."""
- tokenizer: Any = None #: :meta private:
- """Tokenizer."""
- tokenizer_config: Optional[dict] = None
- """
- Configuration parameters specifically for the tokenizer.
- Defaults to an empty dictionary.
- """
- adapter_file: Optional[str] = None
- """
- Path to the adapter file. If provided, applies LoRA layers to the model.
- Defaults to None.
- """
- lazy: bool = False
- """
- If False eval the model parameters to make sure they are
- loaded in memory before returning, otherwise they will be loaded
- when needed. Default: ``False``
- """
- pipeline_kwargs: Optional[dict] = None
- """
- Keyword arguments passed to the pipeline. Defaults include:
- - temp (float): Temperature for generation, default is 0.0.
- - max_tokens (int): Maximum tokens to generate, default is 100.
- - verbose (bool): Whether to output verbose logging, default is False.
- - formatter (Optional[Callable]): A callable to format the output.
- Default is None.
- - repetition_penalty (Optional[float]): The penalty factor for
- repeated sequences, default is None.
- - repetition_context_size (Optional[int]): Size of the context
- for applying repetition penalty, default is None.
- - top_p (float): The cumulative probability threshold for
- top-p filtering, default is 1.0.
-
- """
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @classmethod
- def from_model_id(
- cls,
- model_id: str,
- tokenizer_config: Optional[dict] = None,
- adapter_file: Optional[str] = None,
- lazy: bool = False,
- pipeline_kwargs: Optional[dict] = None,
- **kwargs: Any,
- ) -> MLXPipeline:
- """Construct the pipeline object from model_id and task."""
- try:
- from mlx_lm import load
-
- except ImportError:
- raise ImportError(
- "Could not import mlx_lm python package. "
- "Please install it with `pip install mlx_lm`."
- )
-
- tokenizer_config = tokenizer_config or {}
- if adapter_file:
- model, tokenizer = load(
- model_id, tokenizer_config, adapter_path=adapter_file, lazy=lazy
- )
- else:
- model, tokenizer = load(model_id, tokenizer_config, lazy=lazy)
-
- _pipeline_kwargs = pipeline_kwargs or {}
- return cls(
- model_id=model_id,
- model=model,
- tokenizer=tokenizer,
- tokenizer_config=tokenizer_config,
- adapter_file=adapter_file,
- lazy=lazy,
- pipeline_kwargs=_pipeline_kwargs,
- **kwargs,
- )
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- "model_id": self.model_id,
- "tokenizer_config": self.tokenizer_config,
- "adapter_file": self.adapter_file,
- "lazy": self.lazy,
- "pipeline_kwargs": self.pipeline_kwargs,
- }
-
- @property
- def _llm_type(self) -> str:
- return "mlx_pipeline"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- try:
- from mlx_lm import generate
- from mlx_lm.sample_utils import make_logits_processors, make_sampler
-
- except ImportError:
- raise ImportError(
- "Could not import mlx_lm python package. "
- "Please install it with `pip install mlx_lm`."
- )
-
- pipeline_kwargs = kwargs.get("pipeline_kwargs", self.pipeline_kwargs)
-
- temp: float = pipeline_kwargs.get("temp", 0.0)
- max_tokens: int = pipeline_kwargs.get("max_tokens", 100)
- verbose: bool = pipeline_kwargs.get("verbose", False)
- formatter: Optional[Callable] = pipeline_kwargs.get("formatter", None)
- repetition_penalty: Optional[float] = pipeline_kwargs.get(
- "repetition_penalty", None
- )
- repetition_context_size: Optional[int] = pipeline_kwargs.get(
- "repetition_context_size", None
- )
- top_p: float = pipeline_kwargs.get("top_p", 1.0)
- min_p: float = pipeline_kwargs.get("min_p", 0.0)
- min_tokens_to_keep: int = pipeline_kwargs.get("min_tokens_to_keep", 1)
-
- sampler = make_sampler(temp, top_p, min_p, min_tokens_to_keep)
- logits_processors = make_logits_processors(
- None, repetition_penalty, repetition_context_size
- )
-
- return generate(
- model=self.model,
- tokenizer=self.tokenizer,
- prompt=prompt,
- max_tokens=max_tokens,
- verbose=verbose,
- formatter=formatter,
- sampler=sampler,
- logits_processors=logits_processors,
- )
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- try:
- import mlx.core as mx
- from mlx_lm.sample_utils import make_logits_processors, make_sampler
- from mlx_lm.utils import generate_step
-
- except ImportError:
- raise ImportError(
- "Could not import mlx_lm python package. "
- "Please install it with `pip install mlx_lm`."
- )
-
- pipeline_kwargs = kwargs.get("pipeline_kwargs", self.pipeline_kwargs)
-
- temp: float = pipeline_kwargs.get("temp", 0.0)
- max_new_tokens: int = pipeline_kwargs.get("max_tokens", 100)
- repetition_penalty: Optional[float] = pipeline_kwargs.get(
- "repetition_penalty", None
- )
- repetition_context_size: Optional[int] = pipeline_kwargs.get(
- "repetition_context_size", None
- )
- top_p: float = pipeline_kwargs.get("top_p", 1.0)
- min_p: float = pipeline_kwargs.get("min_p", 0.0)
- min_tokens_to_keep: int = pipeline_kwargs.get("min_tokens_to_keep", 1)
-
- prompt = self.tokenizer.encode(prompt, return_tensors="np")
-
- prompt_tokens = mx.array(prompt[0])
-
- eos_token_id = self.tokenizer.eos_token_id
- detokenizer = self.tokenizer.detokenizer
- detokenizer.reset()
-
- sampler = make_sampler(temp or 0.0, top_p, min_p, min_tokens_to_keep)
-
- logits_processors = make_logits_processors(
- None, repetition_penalty, repetition_context_size
- )
-
- for (token, prob), n in zip(
- generate_step(
- prompt=prompt_tokens,
- model=self.model,
- sampler=sampler,
- logits_processors=logits_processors,
- ),
- range(max_new_tokens),
- ):
- # identify text to yield
- text: Optional[str] = None
- detokenizer.add_token(token)
- detokenizer.finalize()
- text = detokenizer.last_segment
-
- # yield text, if any
- if text:
- chunk = GenerationChunk(text=text)
- if run_manager:
- run_manager.on_llm_new_token(chunk.text)
- yield chunk
-
- # break if stop sequence found
- if token == eos_token_id or (stop is not None and text in stop):
- break
diff --git a/libs/community/langchain_community/llms/modal.py b/libs/community/langchain_community/llms/modal.py
deleted file mode 100644
index a68aa985d3..0000000000
--- a/libs/community/langchain_community/llms/modal.py
+++ /dev/null
@@ -1,101 +0,0 @@
-import logging
-from typing import Any, Dict, List, Mapping, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils.pydantic import get_fields
-from pydantic import ConfigDict, Field, model_validator
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-class Modal(LLM):
- """Modal large language models.
-
- To use, you should have the ``modal-client`` python package installed.
-
- Any parameters that are valid to be passed to the call can be passed
- in, even if not explicitly saved on this class.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import Modal
- modal = Modal(endpoint_url="")
-
- """
-
- endpoint_url: str = ""
- """model endpoint to use"""
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `create` call not
- explicitly specified."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = {field.alias for field in get_fields(cls).values()}
-
- extra = values.get("model_kwargs", {})
- for field_name in list(values):
- if field_name not in all_required_field_names:
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- logger.warning(
- f"""{field_name} was transferred to model_kwargs.
- Please confirm that {field_name} is what you intended."""
- )
- extra[field_name] = values.pop(field_name)
- values["model_kwargs"] = extra
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- **{"endpoint_url": self.endpoint_url},
- **{"model_kwargs": self.model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "modal"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call to Modal endpoint."""
- params = self.model_kwargs or {}
- params = {**params, **kwargs}
- response = requests.post(
- url=self.endpoint_url,
- headers={
- "Content-Type": "application/json",
- },
- json={"prompt": prompt, **params},
- )
- try:
- if prompt in response.json()["prompt"]:
- response_json = response.json()
- except KeyError:
- raise KeyError("LangChain requires 'prompt' key in response.")
- text = response_json["prompt"]
- if stop is not None:
- # I believe this is required since the stop tokens
- # are not enforced by the model parameters
- text = enforce_stop_tokens(text, stop)
- return text
diff --git a/libs/community/langchain_community/llms/moonshot.py b/libs/community/langchain_community/llms/moonshot.py
deleted file mode 100644
index 7f204fa69d..0000000000
--- a/libs/community/langchain_community/llms/moonshot.py
+++ /dev/null
@@ -1,141 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import (
- BaseModel,
- ConfigDict,
- Field,
- SecretStr,
- model_validator,
-)
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-MOONSHOT_SERVICE_URL_BASE = "https://api.moonshot.cn/v1"
-
-
-class _MoonshotClient(BaseModel):
- """An API client that talks to the Moonshot server."""
-
- api_key: SecretStr
- """The API key to use for authentication."""
- base_url: str = MOONSHOT_SERVICE_URL_BASE
-
- def completion(self, request: Any) -> Any:
- headers = {"Authorization": f"Bearer {self.api_key.get_secret_value()}"}
- response = requests.post(
- f"{self.base_url}/chat/completions",
- headers=headers,
- json=request,
- )
- if not response.ok:
- raise ValueError(f"HTTP {response.status_code} error: {response.text}")
- return response.json()["choices"][0]["message"]["content"]
-
-
-class MoonshotCommon(BaseModel):
- """Common parameters for Moonshot LLMs."""
-
- client: Any
- base_url: str = MOONSHOT_SERVICE_URL_BASE
- moonshot_api_key: Optional[SecretStr] = Field(default=None, alias="api_key")
- """Moonshot API key. Get it here: https://platform.moonshot.cn/console/api-keys"""
- model_name: str = Field(default="moonshot-v1-8k", alias="model")
- """Model name. Available models listed here: https://platform.moonshot.cn/pricing"""
- max_tokens: int = 1024
- """Maximum number of tokens to generate."""
- temperature: float = 0.3
- """Temperature parameter (higher values make the model more creative)."""
-
- model_config = ConfigDict(populate_by_name=True, protected_namespaces=())
-
- @property
- def lc_secrets(self) -> dict:
- """A map of constructor argument names to secret ids.
-
- For example,
- {"moonshot_api_key": "MOONSHOT_API_KEY"}
- """
- return {"moonshot_api_key": "MOONSHOT_API_KEY"}
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling OpenAI API."""
- return {
- "model": self.model_name,
- "max_tokens": self.max_tokens,
- "temperature": self.temperature,
- }
-
- @property
- def _invocation_params(self) -> Dict[str, Any]:
- return {**{"model": self.model_name}, **self._default_params}
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra parameters.
- Override the superclass method, prevent the model parameter from being
- overridden.
- """
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- values["moonshot_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(values, "moonshot_api_key", "MOONSHOT_API_KEY")
- )
-
- values["client"] = _MoonshotClient(
- api_key=values["moonshot_api_key"],
- base_url=values["base_url"]
- if "base_url" in values
- else MOONSHOT_SERVICE_URL_BASE,
- )
- return values
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "moonshot"
-
-
-class Moonshot(MoonshotCommon, LLM):
- """Moonshot large language models.
-
- To use, you should have the environment variable ``MOONSHOT_API_KEY`` set with your
- API key. Referenced from https://platform.moonshot.cn/docs
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms.moonshot import Moonshot
-
- moonshot = Moonshot(model="moonshot-v1-8k")
- """
-
- model_config = ConfigDict(
- populate_by_name=True,
- )
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- request = self._invocation_params
- request["messages"] = [{"role": "user", "content": prompt}]
- request.update(kwargs)
- text = self.client.completion(request)
- if stop is not None:
- # This is required since the stop tokens
- # are not enforced by the model parameters
- text = enforce_stop_tokens(text, stop)
-
- return text
diff --git a/libs/community/langchain_community/llms/mosaicml.py b/libs/community/langchain_community/llms/mosaicml.py
deleted file mode 100644
index 15464dc829..0000000000
--- a/libs/community/langchain_community/llms/mosaicml.py
+++ /dev/null
@@ -1,187 +0,0 @@
-from typing import Any, Dict, List, Mapping, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import ConfigDict
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-INSTRUCTION_KEY = "### Instruction:"
-RESPONSE_KEY = "### Response:"
-INTRO_BLURB = (
- "Below is an instruction that describes a task. "
- "Write a response that appropriately completes the request."
-)
-PROMPT_FOR_GENERATION_FORMAT = """{intro}
-{instruction_key}
-{instruction}
-{response_key}
-""".format(
- intro=INTRO_BLURB,
- instruction_key=INSTRUCTION_KEY,
- instruction="{instruction}",
- response_key=RESPONSE_KEY,
-)
-
-
-class MosaicML(LLM):
- """MosaicML LLM service.
-
- To use, you should have the
- environment variable ``MOSAICML_API_TOKEN`` set with your API token, or pass
- it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import MosaicML
- endpoint_url = (
- "https://models.hosted-on.mosaicml.hosting/mpt-7b-instruct/v1/predict"
- )
- mosaic_llm = MosaicML(
- endpoint_url=endpoint_url,
- mosaicml_api_token="my-api-key"
- )
- """
-
- endpoint_url: str = (
- "https://models.hosted-on.mosaicml.hosting/mpt-7b-instruct/v1/predict"
- )
- """Endpoint URL to use."""
- inject_instruction_format: bool = False
- """Whether to inject the instruction format into the prompt."""
- model_kwargs: Optional[dict] = None
- """Keyword arguments to pass to the model."""
- retry_sleep: float = 1.0
- """How long to try sleeping for if a rate limit is encountered"""
-
- mosaicml_api_token: Optional[str] = None
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- mosaicml_api_token = get_from_dict_or_env(
- values, "mosaicml_api_token", "MOSAICML_API_TOKEN"
- )
- values["mosaicml_api_token"] = mosaicml_api_token
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- **{"endpoint_url": self.endpoint_url},
- **{"model_kwargs": _model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "mosaic"
-
- def _transform_prompt(self, prompt: str) -> str:
- """Transform prompt."""
- if self.inject_instruction_format:
- prompt = PROMPT_FOR_GENERATION_FORMAT.format(
- instruction=prompt,
- )
- return prompt
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- is_retry: bool = False,
- **kwargs: Any,
- ) -> str:
- """Call out to a MosaicML LLM inference endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = mosaic_llm.invoke("Tell me a joke.")
- """
- _model_kwargs = self.model_kwargs or {}
-
- prompt = self._transform_prompt(prompt)
-
- payload = {"inputs": [prompt]}
- payload.update(_model_kwargs)
- payload.update(kwargs)
-
- # HTTP headers for authorization
- headers = {
- "Authorization": f"{self.mosaicml_api_token}",
- "Content-Type": "application/json",
- }
-
- # send request
- try:
- response = requests.post(self.endpoint_url, headers=headers, json=payload)
- except requests.exceptions.RequestException as e:
- raise ValueError(f"Error raised by inference endpoint: {e}")
-
- try:
- if response.status_code == 429:
- if not is_retry:
- import time
-
- time.sleep(self.retry_sleep)
-
- return self._call(prompt, stop, run_manager, is_retry=True)
-
- raise ValueError(
- f"Error raised by inference API: rate limit exceeded.\nResponse: "
- f"{response.text}"
- )
-
- parsed_response = response.json()
-
- # The inference API has changed a couple of times, so we add some handling
- # to be robust to multiple response formats.
- if isinstance(parsed_response, dict):
- output_keys = ["data", "output", "outputs"]
- for key in output_keys:
- if key in parsed_response:
- output_item = parsed_response[key]
- break
- else:
- raise ValueError(
- f"No valid key ({', '.join(output_keys)}) in response:"
- f" {parsed_response}"
- )
- if isinstance(output_item, list):
- text = output_item[0]
- else:
- text = output_item
- else:
- raise ValueError(f"Unexpected response type: {parsed_response}")
-
- # Older versions of the API include the input in the output response
- if text.startswith(prompt):
- text = text[len(prompt) :]
-
- except requests.exceptions.JSONDecodeError as e:
- raise ValueError(
- f"Error raised by inference API: {e}.\nResponse: {response.text}"
- )
-
- # TODO: replace when MosaicML supports custom stop tokens natively
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
- return text
diff --git a/libs/community/langchain_community/llms/nlpcloud.py b/libs/community/langchain_community/llms/nlpcloud.py
deleted file mode 100644
index 774122a38f..0000000000
--- a/libs/community/langchain_community/llms/nlpcloud.py
+++ /dev/null
@@ -1,144 +0,0 @@
-from typing import Any, Dict, List, Mapping, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import ConfigDict, SecretStr
-
-
-class NLPCloud(LLM):
- """NLPCloud large language models.
-
- To use, you should have the ``nlpcloud`` python package installed, and the
- environment variable ``NLPCLOUD_API_KEY`` set with your API key.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import NLPCloud
- nlpcloud = NLPCloud(model="finetuned-gpt-neox-20b")
- """
-
- client: Any = None #: :meta private:
- model_name: str = "finetuned-gpt-neox-20b"
- """Model name to use."""
- gpu: bool = True
- """Whether to use a GPU or not"""
- lang: str = "en"
- """Language to use (multilingual addon)"""
- temperature: float = 0.7
- """What sampling temperature to use."""
- max_length: int = 256
- """The maximum number of tokens to generate in the completion."""
- length_no_input: bool = True
- """Whether min_length and max_length should include the length of the input."""
- remove_input: bool = True
- """Remove input text from API response"""
- remove_end_sequence: bool = True
- """Whether or not to remove the end sequence token."""
- bad_words: List[str] = []
- """List of tokens not allowed to be generated."""
- top_p: float = 1.0
- """Total probability mass of tokens to consider at each step."""
- top_k: int = 50
- """The number of highest probability tokens to keep for top-k filtering."""
- repetition_penalty: float = 1.0
- """Penalizes repeated tokens. 1.0 means no penalty."""
- num_beams: int = 1
- """Number of beams for beam search."""
- num_return_sequences: int = 1
- """How many completions to generate for each prompt."""
-
- nlpcloud_api_key: Optional[SecretStr] = None
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- values["nlpcloud_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(values, "nlpcloud_api_key", "NLPCLOUD_API_KEY")
- )
- try:
- import nlpcloud
-
- values["client"] = nlpcloud.Client(
- values["model_name"],
- values["nlpcloud_api_key"].get_secret_value(),
- gpu=values["gpu"],
- lang=values["lang"],
- )
- except ImportError:
- raise ImportError(
- "Could not import nlpcloud python package. "
- "Please install it with `pip install nlpcloud`."
- )
- return values
-
- @property
- def _default_params(self) -> Mapping[str, Any]:
- """Get the default parameters for calling NLPCloud API."""
- return {
- "temperature": self.temperature,
- "max_length": self.max_length,
- "length_no_input": self.length_no_input,
- "remove_input": self.remove_input,
- "remove_end_sequence": self.remove_end_sequence,
- "bad_words": self.bad_words,
- "top_p": self.top_p,
- "top_k": self.top_k,
- "repetition_penalty": self.repetition_penalty,
- "num_beams": self.num_beams,
- "num_return_sequences": self.num_return_sequences,
- }
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- **{"model_name": self.model_name},
- **{"gpu": self.gpu},
- **{"lang": self.lang},
- **self._default_params,
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "nlpcloud"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to NLPCloud's create endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Not supported by this interface (pass in init method)
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = nlpcloud("Tell me a joke.")
- """
- if stop and len(stop) > 1:
- raise ValueError(
- "NLPCloud only supports a single stop sequence per generation."
- "Pass in a list of length 1."
- )
- elif stop and len(stop) == 1:
- end_sequence = stop[0]
- else:
- end_sequence = None
- params = {**self._default_params, **kwargs}
- response = self.client.generation(prompt, end_sequence=end_sequence, **params)
- return response["generated_text"]
diff --git a/libs/community/langchain_community/llms/oci_data_science_model_deployment_endpoint.py b/libs/community/langchain_community/llms/oci_data_science_model_deployment_endpoint.py
deleted file mode 100644
index 83001e6e6e..0000000000
--- a/libs/community/langchain_community/llms/oci_data_science_model_deployment_endpoint.py
+++ /dev/null
@@ -1,970 +0,0 @@
-# Copyright (c) 2023, 2024, Oracle and/or its affiliates.
-
-"""LLM for OCI data science model deployment endpoint."""
-
-import json
-import logging
-import traceback
-from typing import (
- Any,
- AsyncIterator,
- Callable,
- Dict,
- Iterator,
- List,
- Literal,
- Optional,
- Union,
-)
-
-import aiohttp
-import requests
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import BaseLLM, create_base_retry_decorator
-from langchain_core.load.serializable import Serializable
-from langchain_core.outputs import Generation, GenerationChunk, LLMResult
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import Field, model_validator
-
-from langchain_community.utilities.requests import Requests
-
-logger = logging.getLogger(__name__)
-DEFAULT_INFERENCE_ENDPOINT = "/v1/completions"
-
-
-DEFAULT_TIME_OUT = 300
-DEFAULT_CONTENT_TYPE_JSON = "application/json"
-DEFAULT_MODEL_NAME = "odsc-llm"
-
-
-class TokenExpiredError(Exception):
- """Raises when token expired."""
-
-
-class ServerError(Exception):
- """Raises when encounter server error when making inference."""
-
-
-def _create_retry_decorator(
- llm: "BaseOCIModelDeployment",
- *,
- run_manager: Optional[
- Union[AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun]
- ] = None,
-) -> Callable[[Any], Any]:
- """Create a retry decorator."""
- errors = [requests.exceptions.ConnectTimeout, TokenExpiredError]
- decorator = create_base_retry_decorator(
- error_types=errors, max_retries=llm.max_retries, run_manager=run_manager
- )
- return decorator
-
-
-class BaseOCIModelDeployment(Serializable):
- """Base class for LLM deployed on OCI Data Science Model Deployment."""
-
- auth: dict = Field(default_factory=dict, exclude=True)
- """ADS auth dictionary for OCI authentication:
- https://accelerated-data-science.readthedocs.io/en/latest/user_guide/cli/authentication.html.
- This can be generated by calling `ads.common.auth.api_keys()`
- or `ads.common.auth.resource_principal()`. If this is not
- provided then the `ads.common.default_signer()` will be used."""
-
- endpoint: str = ""
- """The uri of the endpoint from the deployed Model Deployment model."""
-
- streaming: bool = False
- """Whether to stream the results or not."""
-
- max_retries: int = 3
- """Maximum number of retries to make when generating."""
-
- default_headers: Optional[Dict[str, Any]] = None
- """The headers to be added to the Model Deployment request."""
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Dict:
- """Checks if oracle-ads is installed and
- get credentials/endpoint from environment.
- """
- try:
- import ads
-
- except ImportError as ex:
- raise ImportError(
- "Could not import ads python package. "
- "Please install it with `pip install oracle_ads`."
- ) from ex
-
- if not values.get("auth", None):
- values["auth"] = ads.common.auth.default_signer()
-
- values["endpoint"] = get_from_dict_or_env(
- values,
- "endpoint",
- "OCI_LLM_ENDPOINT",
- )
- return values
-
- def _headers(
- self, is_async: Optional[bool] = False, body: Optional[dict] = None
- ) -> Dict:
- """Construct and return the headers for a request.
-
- Args:
- is_async (bool, optional): Indicates if the request is asynchronous.
- Defaults to `False`.
- body (optional): The request body to be included in the headers if
- the request is asynchronous.
-
- Returns:
- Dict: A dictionary containing the appropriate headers for the request.
- """
- headers = self.default_headers or {}
- if is_async:
- signer = self.auth["signer"]
- _req = requests.Request("POST", self.endpoint, json=body)
- req = _req.prepare()
- req = signer(req)
- for key, value in req.headers.items():
- headers[key] = value
-
- if self.streaming:
- headers.update(
- {
- "enable-streaming": "true",
- "Accept": "text/event-stream",
- }
- )
- return headers
-
- headers.update(
- {
- "Content-Type": DEFAULT_CONTENT_TYPE_JSON,
- "enable-streaming": "true",
- "Accept": "text/event-stream",
- }
- if self.streaming
- else {
- "Content-Type": DEFAULT_CONTENT_TYPE_JSON,
- }
- )
-
- return headers
-
- def completion_with_retry(
- self, run_manager: Optional[CallbackManagerForLLMRun] = None, **kwargs: Any
- ) -> Any:
- """Use tenacity to retry the completion call."""
- retry_decorator = _create_retry_decorator(self, run_manager=run_manager)
-
- @retry_decorator
- def _completion_with_retry(**kwargs: Any) -> Any:
- try:
- request_timeout = kwargs.pop("request_timeout", DEFAULT_TIME_OUT)
- data = kwargs.pop("data")
- stream = kwargs.pop("stream", self.streaming)
-
- request = Requests(
- headers=self._headers(), auth=self.auth.get("signer")
- )
- response = request.post(
- url=self.endpoint,
- data=data,
- timeout=request_timeout,
- stream=stream,
- **kwargs,
- )
- self._check_response(response)
- return response
- except TokenExpiredError as e:
- raise e
- except Exception as err:
- traceback.print_exc()
- logger.debug(
- f"Requests payload: {data}. Requests arguments: "
- f"url={self.endpoint},timeout={request_timeout},stream={stream}. "
- f"Additional request kwargs={kwargs}."
- )
- raise RuntimeError(
- f"Error occurs by inference endpoint: {str(err)}"
- ) from err
-
- return _completion_with_retry(**kwargs)
-
- async def acompletion_with_retry(
- self,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Any:
- """Use tenacity to retry the async completion call."""
- retry_decorator = _create_retry_decorator(self, run_manager=run_manager)
-
- @retry_decorator
- async def _completion_with_retry(**kwargs: Any) -> Any:
- try:
- request_timeout = kwargs.pop("request_timeout", DEFAULT_TIME_OUT)
- data = kwargs.pop("data")
- stream = kwargs.pop("stream", self.streaming)
-
- request = Requests(headers=self._headers(is_async=True, body=data))
- if stream:
- response = request.apost(
- url=self.endpoint,
- data=data,
- timeout=request_timeout,
- )
- return self._aiter_sse(response)
- else:
- async with request.apost(
- url=self.endpoint,
- data=data,
- timeout=request_timeout,
- ) as resp:
- self._check_response(resp)
- data = await resp.json()
- return data
- except TokenExpiredError as e:
- raise e
- except Exception as err:
- traceback.print_exc()
- logger.debug(
- f"Requests payload: `{data}`. "
- f"Stream mode={stream}. "
- f"Requests kwargs: url={self.endpoint}, timeout={request_timeout}."
- )
- raise RuntimeError(
- f"Error occurs by inference endpoint: {str(err)}"
- ) from err
-
- return await _completion_with_retry(**kwargs)
-
- def _check_response(self, response: Any) -> None:
- """Handle server error by checking the response status.
-
- Args:
- response:
- The response object from either `requests` or `aiohttp` library.
-
- Raises:
- TokenExpiredError:
- If the response status code is 401 and the token refresh is successful.
- ServerError:
- If any other HTTP error occurs.
- """
- try:
- response.raise_for_status()
- except requests.exceptions.HTTPError as http_err:
- status_code = (
- response.status_code
- if hasattr(response, "status_code")
- else response.status
- )
- if status_code == 401 and self._refresh_signer():
- raise TokenExpiredError() from http_err
-
- raise ServerError(
- f"Server error: {str(http_err)}. \nMessage: {response.text}"
- ) from http_err
-
- def _parse_stream(self, lines: Iterator[bytes]) -> Iterator[str]:
- """Parse a stream of byte lines and yield parsed string lines.
-
- Args:
- lines (Iterator[bytes]):
- An iterator that yields lines in byte format.
-
- Yields:
- Iterator[str]:
- An iterator that yields parsed lines as strings.
- """
- for line in lines:
- _line = self._parse_stream_line(line)
- if _line is not None:
- yield _line
-
- async def _parse_stream_async(
- self,
- lines: aiohttp.StreamReader,
- ) -> AsyncIterator[str]:
- """
- Asynchronously parse a stream of byte lines and yield parsed string lines.
-
- Args:
- lines (aiohttp.StreamReader):
- An `aiohttp.StreamReader` object that yields lines in byte format.
-
- Yields:
- AsyncIterator[str]:
- An asynchronous iterator that yields parsed lines as strings.
- """
- async for line in lines:
- _line = self._parse_stream_line(line)
- if _line is not None:
- yield _line
-
- def _parse_stream_line(self, line: bytes) -> Optional[str]:
- """Parse a single byte line and return a processed string line if valid.
-
- Args:
- line (bytes): A single line in byte format.
-
- Returns:
- Optional[str]:
- The processed line as a string if valid, otherwise `None`.
- """
- line = line.strip()
- if not line:
- return None
- _line = line.decode("utf-8")
-
- if _line.lower().startswith("data:"):
- _line = _line[5:].lstrip()
-
- if _line.startswith("[DONE]"):
- return None
- return _line
- return None
-
- async def _aiter_sse(
- self,
- async_cntx_mgr: Any,
- ) -> AsyncIterator[str]:
- """Asynchronously iterate over server-sent events (SSE).
-
- Args:
- async_cntx_mgr: An asynchronous context manager that yields a client
- response object.
-
- Yields:
- AsyncIterator[str]: An asynchronous iterator that yields parsed server-sent
- event lines as json string.
- """
- async with async_cntx_mgr as client_resp:
- self._check_response(client_resp)
- async for line in self._parse_stream_async(client_resp.content):
- yield line
-
- def _refresh_signer(self) -> bool:
- """Attempt to refresh the security token using the signer.
-
- Returns:
- bool: `True` if the token was successfully refreshed, `False` otherwise.
- """
- if self.auth.get("signer", None) and hasattr(
- self.auth["signer"], "refresh_security_token"
- ):
- self.auth["signer"].refresh_security_token()
- return True
- return False
-
- @classmethod
- def is_lc_serializable(cls) -> bool:
- """Return whether this model can be serialized by LangChain."""
- return True
-
-
-class OCIModelDeploymentLLM(BaseLLM, BaseOCIModelDeployment):
- """LLM deployed on OCI Data Science Model Deployment.
-
- To use, you must provide the model HTTP endpoint from your deployed
- model, e.g. https://modeldeployment..oci.customer-oci.com//predict.
-
- To authenticate, `oracle-ads` has been used to automatically load
- credentials: https://accelerated-data-science.readthedocs.io/en/latest/user_guide/cli/authentication.html
-
- Make sure to have the required policies to access the OCI Data
- Science Model Deployment endpoint. See:
- https://docs.oracle.com/en-us/iaas/data-science/using/model-dep-policies-auth.htm#model_dep_policies_auth__predict-endpoint
-
- Example:
-
- .. code-block:: python
-
- from langchain_community.llms import OCIModelDeploymentLLM
-
- llm = OCIModelDeploymentLLM(
- endpoint="https://modeldeployment.us-ashburn-1.oci.customer-oci.com//predict",
- model="odsc-llm",
- streaming=True,
- model_kwargs={"frequency_penalty": 1.0},
- headers={
- "route": "/v1/completions",
- # other request headers ...
- }
- )
- llm.invoke("tell me a joke.")
-
- Customized Usage:
-
- User can inherit from our base class and overrwrite the `_process_response`, `_process_stream_response`,
- `_construct_json_body` for satisfying customized needed.
-
- .. code-block:: python
-
- from langchain_community.llms import OCIModelDeploymentLLM
-
- class MyCutomizedModel(OCIModelDeploymentLLM):
- def _process_stream_response(self, response_json:dict) -> GenerationChunk:
- print("My customized output stream handler.")
- return GenerationChunk()
-
- def _process_response(self, response_json:dict) -> List[Generation]:
- print("My customized output handler.")
- return [Generation()]
-
- def _construct_json_body(self, prompt: str, param:dict) -> dict:
- print("My customized input handler.")
- return {}
-
- llm = MyCutomizedModel(
- endpoint=f"https://modeldeployment.us-ashburn-1.oci.customer-oci.com/{ocid}/predict",
- model="",
- }
-
- llm.invoke("tell me a joke.")
-
- """ # noqa: E501
-
- model: str = DEFAULT_MODEL_NAME
- """The name of the model."""
-
- max_tokens: int = 256
- """Denotes the number of tokens to predict per generation."""
-
- temperature: float = 0.2
- """A non-negative float that tunes the degree of randomness in generation."""
-
- k: int = 50
- """Number of most likely tokens to consider at each step."""
-
- p: float = 0.75
- """Total probability mass of tokens to consider at each step."""
-
- best_of: int = 1
- """Generates best_of completions server-side and returns the "best"
- (the one with the highest log probability per token).
- """
-
- stop: Optional[List[str]] = None
- """Stop words to use when generating. Model output is cut off
- at the first occurrence of any of these substrings."""
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Keyword arguments to pass to the model."""
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "oci_model_deployment_endpoint"
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters."""
- return {
- "best_of": self.best_of,
- "max_tokens": self.max_tokens,
- "model": self.model,
- "stop": self.stop,
- "stream": self.streaming,
- "temperature": self.temperature,
- "top_k": self.k,
- "top_p": self.p,
- }
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- **{"endpoint": self.endpoint, "model_kwargs": _model_kwargs},
- **self._default_params,
- }
-
- def _headers(
- self, is_async: Optional[bool] = False, body: Optional[dict] = None
- ) -> Dict:
- """Construct and return the headers for a request.
-
- Args:
- is_async (bool, optional): Indicates if the request is asynchronous.
- Defaults to `False`.
- body (optional): The request body to be included in the headers if
- the request is asynchronous.
-
- Returns:
- Dict: A dictionary containing the appropriate headers for the request.
- """
- return {
- "route": DEFAULT_INFERENCE_ENDPOINT,
- **super()._headers(is_async=is_async, body=body),
- }
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call out to OCI Data Science Model Deployment endpoint with k unique prompts.
-
- Args:
- prompts: The prompts to pass into the service.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The full LLM output.
-
- Example:
- .. code-block:: python
-
- response = llm.invoke("Tell me a joke.")
- response = llm.generate(["Tell me a joke."])
- """
- generations: List[List[Generation]] = []
- params = self._invocation_params(stop, **kwargs)
- for prompt in prompts:
- body = self._construct_json_body(prompt, params)
- if self.streaming:
- generation = GenerationChunk(text="")
- for chunk in self._stream(
- prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- generation += chunk
- generations.append([generation])
- else:
- res = self.completion_with_retry(
- data=body,
- run_manager=run_manager,
- **kwargs,
- )
- generations.append(self._process_response(res.json()))
- return LLMResult(generations=generations)
-
- async def _agenerate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call out to OCI Data Science Model Deployment endpoint async with k unique prompts.
-
- Args:
- prompts: The prompts to pass into the service.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The full LLM output.
-
- Example:
- .. code-block:: python
-
- response = await llm.ainvoke("Tell me a joke.")
- response = await llm.agenerate(["Tell me a joke."])
- """ # noqa: E501
- generations: List[List[Generation]] = []
- params = self._invocation_params(stop, **kwargs)
- for prompt in prompts:
- body = self._construct_json_body(prompt, params)
- if self.streaming:
- generation = GenerationChunk(text="")
- async for chunk in self._astream(
- prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- generation += chunk
- generations.append([generation])
- else:
- res = await self.acompletion_with_retry(
- data=body,
- run_manager=run_manager,
- **kwargs,
- )
- generations.append(self._process_response(res))
- return LLMResult(generations=generations)
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- """Stream OCI Data Science Model Deployment endpoint on given prompt.
-
-
- Args:
- prompt (str):
- The prompt to pass into the model.
- stop (List[str], Optional):
- List of stop words to use when generating.
- kwargs:
- requests_kwargs:
- Additional ``**kwargs`` to pass to requests.post
-
- Returns:
- An iterator of GenerationChunks.
-
-
- Example:
-
- .. code-block:: python
-
- response = llm.stream("Tell me a joke.")
-
- """
- requests_kwargs = kwargs.pop("requests_kwargs", {})
- self.streaming = True
- params = self._invocation_params(stop, **kwargs)
- body = self._construct_json_body(prompt, params)
-
- response = self.completion_with_retry(
- data=body, run_manager=run_manager, stream=True, **requests_kwargs
- )
- for line in self._parse_stream(response.iter_lines()):
- chunk = self._handle_sse_line(line)
- if run_manager:
- run_manager.on_llm_new_token(chunk.text, chunk=chunk)
-
- yield chunk
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- """Stream OCI Data Science Model Deployment endpoint async on given prompt.
-
-
- Args:
- prompt (str):
- The prompt to pass into the model.
- stop (List[str], Optional):
- List of stop words to use when generating.
- kwargs:
- requests_kwargs:
- Additional ``**kwargs`` to pass to requests.post
-
- Returns:
- An iterator of GenerationChunks.
-
-
- Example:
-
- .. code-block:: python
-
- async for chunk in llm.astream(("Tell me a joke."):
- print(chunk, end="", flush=True)
-
- """
- requests_kwargs = kwargs.pop("requests_kwargs", {})
- self.streaming = True
- params = self._invocation_params(stop, **kwargs)
- body = self._construct_json_body(prompt, params)
-
- async for line in await self.acompletion_with_retry(
- data=body, run_manager=run_manager, stream=True, **requests_kwargs
- ):
- chunk = self._handle_sse_line(line)
- if run_manager:
- await run_manager.on_llm_new_token(chunk.text, chunk=chunk)
- yield chunk
-
- def _construct_json_body(self, prompt: str, params: dict) -> dict:
- """Constructs the request body as a dictionary (JSON)."""
- return {
- "prompt": prompt,
- **params,
- }
-
- def _invocation_params(
- self, stop: Optional[List[str]] = None, **kwargs: Any
- ) -> dict:
- """Combines the invocation parameters with default parameters."""
- params = self._default_params
- _model_kwargs = self.model_kwargs or {}
- params["stop"] = stop or params.get("stop", [])
- return {**params, **_model_kwargs, **kwargs}
-
- def _process_stream_response(self, response_json: dict) -> GenerationChunk:
- """Formats streaming response for OpenAI spec into GenerationChunk."""
- try:
- choice = response_json["choices"][0]
- if not isinstance(choice, dict):
- raise TypeError("Endpoint response is not well formed.")
- except (KeyError, IndexError, TypeError) as e:
- raise ValueError("Error while formatting response payload.") from e
-
- return GenerationChunk(text=choice.get("text", ""))
-
- def _process_response(self, response_json: dict) -> List[Generation]:
- """Formats response in OpenAI spec.
-
- Args:
- response_json (dict): The JSON response from the chat model endpoint.
-
- Returns:
- ChatResult: An object containing the list of `ChatGeneration` objects
- and additional LLM output information.
-
- Raises:
- ValueError: If the response JSON is not well-formed or does not
- contain the expected structure.
-
- """
- generations = []
- try:
- choices = response_json["choices"]
- if not isinstance(choices, list):
- raise TypeError("Endpoint response is not well formed.")
- except (KeyError, TypeError) as e:
- raise ValueError("Error while formatting response payload.") from e
-
- for choice in choices:
- gen = Generation(
- text=choice.get("text"),
- generation_info=self._generate_info(choice),
- )
- generations.append(gen)
-
- return generations
-
- def _generate_info(self, choice: dict) -> Any:
- """Extracts generation info from the response."""
- gen_info = {}
- finish_reason = choice.get("finish_reason", None)
- logprobs = choice.get("logprobs", None)
- index = choice.get("index", None)
- if finish_reason:
- gen_info.update({"finish_reason": finish_reason})
- if logprobs is not None:
- gen_info.update({"logprobs": logprobs})
- if index is not None:
- gen_info.update({"index": index})
-
- return gen_info or None
-
- def _handle_sse_line(self, line: str) -> GenerationChunk:
- try:
- obj = json.loads(line)
- return self._process_stream_response(obj)
- except Exception:
- return GenerationChunk(text="")
-
-
-class OCIModelDeploymentTGI(OCIModelDeploymentLLM):
- """OCI Data Science Model Deployment TGI Endpoint.
-
- To use, you must provide the model HTTP endpoint from your deployed
- model, e.g. https://modeldeployment..oci.customer-oci.com//predict.
-
- To authenticate, `oracle-ads` has been used to automatically load
- credentials: https://accelerated-data-science.readthedocs.io/en/latest/user_guide/cli/authentication.html
-
- Make sure to have the required policies to access the OCI Data
- Science Model Deployment endpoint. See:
- https://docs.oracle.com/en-us/iaas/data-science/using/model-dep-policies-auth.htm#model_dep_policies_auth__predict-endpoint
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import OCIModelDeploymentTGI
-
- llm = OCIModelDeploymentTGI(
- endpoint="https://modeldeployment..oci.customer-oci.com//predict",
- api="/v1/completions",
- streaming=True,
- temperature=0.2,
- seed=42,
- # other model parameters ...
- )
-
- """
-
- api: Literal["/generate", "/v1/completions"] = "/v1/completions"
- """Api spec."""
-
- frequency_penalty: float = 0.0
- """Penalizes repeated tokens according to frequency. Between 0 and 1."""
-
- seed: Optional[int] = None
- """Random sampling seed"""
-
- repetition_penalty: Optional[float] = None
- """The parameter for repetition penalty. 1.0 means no penalty."""
-
- suffix: Optional[str] = None
- """The text to append to the prompt. """
-
- do_sample: bool = True
- """If set to True, this parameter enables decoding strategies such as
- multi-nominal sampling, beam-search multi-nominal sampling, Top-K
- sampling and Top-p sampling.
- """
-
- watermark: bool = True
- """Watermarking with `A Watermark for Large Language Models `_.
- Defaults to True."""
-
- return_full_text: bool = False
- """Whether to prepend the prompt to the generated text. Defaults to False."""
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "oci_model_deployment_tgi_endpoint"
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for invoking OCI model deployment TGI endpoint."""
- return (
- {
- "model": self.model, # can be any
- "frequency_penalty": self.frequency_penalty,
- "max_tokens": self.max_tokens,
- "repetition_penalty": self.repetition_penalty,
- "temperature": self.temperature,
- "top_p": self.p,
- "seed": self.seed,
- "stream": self.streaming,
- "suffix": self.suffix,
- "stop": self.stop,
- }
- if self.api == "/v1/completions"
- else {
- "best_of": self.best_of,
- "max_new_tokens": self.max_tokens,
- "temperature": self.temperature,
- "top_k": (
- self.k if self.k > 0 else None
- ), # `top_k` must be strictly positive'
- "top_p": self.p,
- "do_sample": self.do_sample,
- "return_full_text": self.return_full_text,
- "watermark": self.watermark,
- "stop": self.stop,
- }
- )
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- **{
- "endpoint": self.endpoint,
- "api": self.api,
- "model_kwargs": _model_kwargs,
- },
- **self._default_params,
- }
-
- def _construct_json_body(self, prompt: str, params: dict) -> dict:
- """Construct request payload."""
- if self.api == "/v1/completions":
- return super()._construct_json_body(prompt, params)
-
- return {
- "inputs": prompt,
- "parameters": params,
- }
-
- def _process_response(self, response_json: dict) -> List[Generation]:
- """Formats response."""
- if self.api == "/v1/completions":
- return super()._process_response(response_json)
-
- try:
- text = response_json["generated_text"]
- except KeyError as e:
- raise ValueError(
- f"Error while formatting response payload.response_json={response_json}"
- ) from e
-
- return [Generation(text=text)]
-
-
-class OCIModelDeploymentVLLM(OCIModelDeploymentLLM):
- """VLLM deployed on OCI Data Science Model Deployment
-
- To use, you must provide the model HTTP endpoint from your deployed
- model, e.g. https://modeldeployment..oci.customer-oci.com//predict.
-
- To authenticate, `oracle-ads` has been used to automatically load
- credentials: https://accelerated-data-science.readthedocs.io/en/latest/user_guide/cli/authentication.html
-
- Make sure to have the required policies to access the OCI Data
- Science Model Deployment endpoint. See:
- https://docs.oracle.com/en-us/iaas/data-science/using/model-dep-policies-auth.htm#model_dep_policies_auth__predict-endpoint
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import OCIModelDeploymentVLLM
-
- llm = OCIModelDeploymentVLLM(
- endpoint="https://modeldeployment..oci.customer-oci.com//predict",
- model="odsc-llm",
- streaming=False,
- temperature=0.2,
- max_tokens=512,
- n=3,
- best_of=3,
- # other model parameters
- )
-
- """
-
- n: int = 1
- """Number of output sequences to return for the given prompt."""
-
- k: int = -1
- """Number of most likely tokens to consider at each step."""
-
- frequency_penalty: float = 0.0
- """Penalizes repeated tokens according to frequency. Between 0 and 1."""
-
- presence_penalty: float = 0.0
- """Penalizes repeated tokens. Between 0 and 1."""
-
- use_beam_search: bool = False
- """Whether to use beam search instead of sampling."""
-
- ignore_eos: bool = False
- """Whether to ignore the EOS token and continue generating tokens after
- the EOS token is generated."""
-
- logprobs: Optional[int] = None
- """Number of log probabilities to return per output token."""
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "oci_model_deployment_vllm_endpoint"
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling vllm."""
- return {
- "best_of": self.best_of,
- "frequency_penalty": self.frequency_penalty,
- "ignore_eos": self.ignore_eos,
- "logprobs": self.logprobs,
- "max_tokens": self.max_tokens,
- "model": self.model,
- "n": self.n,
- "presence_penalty": self.presence_penalty,
- "stop": self.stop,
- "stream": self.streaming,
- "temperature": self.temperature,
- "top_k": self.k,
- "top_p": self.p,
- "use_beam_search": self.use_beam_search,
- }
diff --git a/libs/community/langchain_community/llms/oci_generative_ai.py b/libs/community/langchain_community/llms/oci_generative_ai.py
deleted file mode 100644
index a2d5dc8b84..0000000000
--- a/libs/community/langchain_community/llms/oci_generative_ai.py
+++ /dev/null
@@ -1,382 +0,0 @@
-from __future__ import annotations
-
-import json
-from abc import ABC, abstractmethod
-from enum import Enum
-from typing import Any, Dict, Iterator, List, Mapping, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import pre_init
-from pydantic import BaseModel, ConfigDict, Field
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-CUSTOM_ENDPOINT_PREFIX = "ocid1.generativeaiendpoint"
-
-
-class Provider(ABC):
- @property
- @abstractmethod
- def stop_sequence_key(self) -> str: ...
-
- @abstractmethod
- def completion_response_to_text(self, response: Any) -> str: ...
-
-
-class CohereProvider(Provider):
- stop_sequence_key: str = "stop_sequences"
-
- def __init__(self) -> None:
- from oci.generative_ai_inference import models
-
- self.llm_inference_request = models.CohereLlmInferenceRequest
-
- def completion_response_to_text(self, response: Any) -> str:
- return response.data.inference_response.generated_texts[0].text
-
-
-class MetaProvider(Provider):
- stop_sequence_key: str = "stop"
-
- def __init__(self) -> None:
- from oci.generative_ai_inference import models
-
- self.llm_inference_request = models.LlamaLlmInferenceRequest
-
- def completion_response_to_text(self, response: Any) -> str:
- return response.data.inference_response.choices[0].text
-
-
-class OCIAuthType(Enum):
- """OCI authentication types as enumerator."""
-
- API_KEY = 1
- SECURITY_TOKEN = 2
- INSTANCE_PRINCIPAL = 3
- RESOURCE_PRINCIPAL = 4
-
-
-class OCIGenAIBase(BaseModel, ABC):
- """Base class for OCI GenAI models"""
-
- client: Any = Field(default=None, exclude=True) #: :meta private:
-
- auth_type: Optional[str] = "API_KEY"
- """Authentication type, could be
-
- API_KEY,
- SECURITY_TOKEN,
- INSTANCE_PRINCIPAL,
- RESOURCE_PRINCIPAL
-
- If not specified, API_KEY will be used
- """
-
- auth_profile: Optional[str] = "DEFAULT"
- """The name of the profile in ~/.oci/config
- If not specified , DEFAULT will be used
- """
-
- auth_file_location: Optional[str] = "~/.oci/config"
- """Path to the config file.
- If not specified, ~/.oci/config will be used
- """
-
- model_id: Optional[str] = None
- """Id of the model to call, e.g., cohere.command"""
-
- provider: Optional[str] = None
- """Provider name of the model. Default to None,
- will try to be derived from the model_id
- otherwise, requires user input
- """
-
- model_kwargs: Optional[Dict] = None
- """Keyword arguments to pass to the model"""
-
- service_endpoint: Optional[str] = None
- """service endpoint url"""
-
- compartment_id: Optional[str] = None
- """OCID of compartment"""
-
- is_stream: bool = False
- """Whether to stream back partial progress"""
-
- model_config = ConfigDict(
- extra="forbid", arbitrary_types_allowed=True, protected_namespaces=()
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that OCI config and python package exists in environment."""
-
- # Skip creating new client if passed in constructor
- if values["client"] is not None:
- return values
-
- try:
- import oci
-
- client_kwargs = {
- "config": {},
- "signer": None,
- "service_endpoint": values["service_endpoint"],
- "retry_strategy": oci.retry.DEFAULT_RETRY_STRATEGY,
- "timeout": (10, 240), # default timeout config for OCI Gen AI service
- }
-
- if values["auth_type"] == OCIAuthType(1).name:
- client_kwargs["config"] = oci.config.from_file(
- file_location=values["auth_file_location"],
- profile_name=values["auth_profile"],
- )
- client_kwargs.pop("signer", None)
- elif values["auth_type"] == OCIAuthType(2).name:
-
- def make_security_token_signer(oci_config): # type: ignore[no-untyped-def]
- pk = oci.signer.load_private_key_from_file(
- oci_config.get("key_file"), None
- )
- with open(
- oci_config.get("security_token_file"), encoding="utf-8"
- ) as f:
- st_string = f.read()
- return oci.auth.signers.SecurityTokenSigner(st_string, pk)
-
- client_kwargs["config"] = oci.config.from_file(
- file_location=values["auth_file_location"],
- profile_name=values["auth_profile"],
- )
- client_kwargs["signer"] = make_security_token_signer(
- oci_config=client_kwargs["config"]
- )
- elif values["auth_type"] == OCIAuthType(3).name:
- client_kwargs["signer"] = (
- oci.auth.signers.InstancePrincipalsSecurityTokenSigner()
- )
- elif values["auth_type"] == OCIAuthType(4).name:
- client_kwargs["signer"] = (
- oci.auth.signers.get_resource_principals_signer()
- )
- else:
- raise ValueError(
- "Please provide valid value to auth_type, "
- f"{values['auth_type']} is not valid."
- )
-
- values["client"] = oci.generative_ai_inference.GenerativeAiInferenceClient(
- **client_kwargs
- )
-
- except ImportError as ex:
- raise ModuleNotFoundError(
- "Could not import oci python package. "
- "Please make sure you have the oci package installed."
- ) from ex
- except Exception as e:
- raise ValueError(
- """Could not authenticate with OCI client.
- If INSTANCE_PRINCIPAL or RESOURCE_PRINCIPAL is used,
- please check the specified
- auth_profile, auth_file_location and auth_type are valid.""",
- e,
- ) from e
-
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- **{"model_kwargs": _model_kwargs},
- }
-
- def _get_provider(self, provider_map: Mapping[str, Any]) -> Any:
- if self.provider is not None:
- provider = self.provider
- else:
- if self.model_id is None:
- raise ValueError(
- "model_id is required to derive the provider, "
- "please provide the provider explicitly or specify "
- "the model_id to derive the provider."
- )
- provider = self.model_id.split(".")[0].lower()
-
- if provider not in provider_map:
- raise ValueError(
- f"Invalid provider derived from model_id: {self.model_id} "
- "Please explicitly pass in the supported provider "
- "when using custom endpoint"
- )
- return provider_map[provider]
-
-
-class OCIGenAI(LLM, OCIGenAIBase):
- """OCI large language models.
-
- To authenticate, the OCI client uses the methods described in
- https://docs.oracle.com/en-us/iaas/Content/API/Concepts/sdk_authentication_methods.htm
-
- The authentifcation method is passed through auth_type and should be one of:
- API_KEY (default), SECURITY_TOKEN, INSTANCE_PRINCIPAL, RESOURCE_PRINCIPAL
-
- Make sure you have the required policies (profile/roles) to
- access the OCI Generative AI service.
- If a specific config profile is used, you must pass
- the name of the profile (from ~/.oci/config) through auth_profile.
- If a specific config file location is used, you must pass
- the file location where profile name configs present
- through auth_file_location
-
- To use, you must provide the compartment id
- along with the endpoint url, and model id
- as named parameters to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import OCIGenAI
-
- llm = OCIGenAI(
- model_id="MY_MODEL_ID",
- service_endpoint="https://inference.generativeai.us-chicago-1.oci.oraclecloud.com",
- compartment_id="MY_OCID"
- )
- """
-
- model_config = ConfigDict(
- extra="forbid",
- arbitrary_types_allowed=True,
- )
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "oci_generative_ai_completion"
-
- @property
- def _provider_map(self) -> Mapping[str, Any]:
- """Get the provider map"""
- return {
- "cohere": CohereProvider(),
- "meta": MetaProvider(),
- }
-
- @property
- def _provider(self) -> Any:
- """Get the internal provider object"""
- return self._get_provider(provider_map=self._provider_map)
-
- def _prepare_invocation_object(
- self, prompt: str, stop: Optional[List[str]], kwargs: Dict[str, Any]
- ) -> Dict[str, Any]:
- from oci.generative_ai_inference import models
-
- _model_kwargs = self.model_kwargs or {}
- if stop is not None:
- _model_kwargs[self._provider.stop_sequence_key] = stop
-
- if self.model_id is None:
- raise ValueError(
- "model_id is required to call the model, please provide the model_id."
- )
-
- if self.model_id.startswith(CUSTOM_ENDPOINT_PREFIX):
- serving_mode = models.DedicatedServingMode(endpoint_id=self.model_id)
- else:
- serving_mode = models.OnDemandServingMode(model_id=self.model_id)
-
- inference_params = {**_model_kwargs, **kwargs}
- inference_params["prompt"] = prompt
- inference_params["is_stream"] = self.is_stream
-
- invocation_obj = models.GenerateTextDetails(
- compartment_id=self.compartment_id,
- serving_mode=serving_mode,
- inference_request=self._provider.llm_inference_request(**inference_params),
- )
-
- return invocation_obj
-
- def _process_response(self, response: Any, stop: Optional[List[str]]) -> str:
- text = self._provider.completion_response_to_text(response)
-
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
-
- return text
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to OCIGenAI generate endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = llm.invoke("Tell me a joke.")
- """
- if self.is_stream:
- text = ""
- for chunk in self._stream(prompt, stop, run_manager, **kwargs):
- text += chunk.text
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
- return text
-
- invocation_obj = self._prepare_invocation_object(prompt, stop, kwargs)
- response = self.client.generate_text(invocation_obj)
- return self._process_response(response, stop)
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- """Stream OCIGenAI LLM on given prompt.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- An iterator of GenerationChunks.
-
- Example:
- .. code-block:: python
-
- response = llm.stream("Tell me a joke.")
- """
-
- self.is_stream = True
- invocation_obj = self._prepare_invocation_object(prompt, stop, kwargs)
- response = self.client.generate_text(invocation_obj)
-
- for event in response.data.events():
- json_load = json.loads(event.data)
- if "text" in json_load:
- event_data_text = json_load["text"]
- else:
- event_data_text = ""
- chunk = GenerationChunk(text=event_data_text)
- if run_manager:
- run_manager.on_llm_new_token(chunk.text, chunk=chunk)
- yield chunk
diff --git a/libs/community/langchain_community/llms/octoai_endpoint.py b/libs/community/langchain_community/llms/octoai_endpoint.py
deleted file mode 100644
index ef519735b7..0000000000
--- a/libs/community/langchain_community/llms/octoai_endpoint.py
+++ /dev/null
@@ -1,117 +0,0 @@
-from typing import Any, Dict
-
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import Field, SecretStr
-
-from langchain_community.llms.openai import BaseOpenAI
-from langchain_community.utils.openai import is_openai_v1
-
-DEFAULT_BASE_URL = "https://text.octoai.run/v1/"
-DEFAULT_MODEL = "codellama-7b-instruct"
-
-
-class OctoAIEndpoint(BaseOpenAI):
- """OctoAI LLM Endpoints - OpenAI compatible.
-
- OctoAIEndpoint is a class to interact with OctoAI Compute Service large
- language model endpoints.
-
- To use, you should have the environment variable ``OCTOAI_API_TOKEN`` set
- with your API token, or pass it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms.octoai_endpoint import OctoAIEndpoint
-
- llm = OctoAIEndpoint(
- model="llama-2-13b-chat-fp16",
- max_tokens=200,
- presence_penalty=0,
- temperature=0.1,
- top_p=0.9,
- )
-
- """
-
- """Key word arguments to pass to the model."""
- octoai_api_base: str = Field(default=DEFAULT_BASE_URL)
- octoai_api_token: SecretStr = Field(default=SecretStr(""))
- model_name: str = Field(default=DEFAULT_MODEL)
-
- @classmethod
- def is_lc_serializable(cls) -> bool:
- return False
-
- @property
- def _invocation_params(self) -> Dict[str, Any]:
- """Get the parameters used to invoke the model."""
-
- params: Dict[str, Any] = {
- "model": self.model_name,
- **self._default_params,
- }
- if not is_openai_v1():
- params.update(
- {
- "api_key": self.octoai_api_token.get_secret_value(),
- "api_base": self.octoai_api_base,
- }
- )
-
- return {**params, **super()._invocation_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "octoai_endpoint"
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- values["octoai_api_base"] = get_from_dict_or_env(
- values,
- "octoai_api_base",
- "OCTOAI_API_BASE",
- default=DEFAULT_BASE_URL,
- )
- values["octoai_api_token"] = convert_to_secret_str(
- get_from_dict_or_env(values, "octoai_api_token", "OCTOAI_API_TOKEN")
- )
- values["model_name"] = get_from_dict_or_env(
- values,
- "model_name",
- "MODEL_NAME",
- default=DEFAULT_MODEL,
- )
-
- try:
- import openai
-
- if is_openai_v1():
- client_params = {
- "api_key": values["octoai_api_token"].get_secret_value(),
- "base_url": values["octoai_api_base"],
- }
- if not values.get("client"):
- values["client"] = openai.OpenAI(**client_params).completions
- if not values.get("async_client"):
- values["async_client"] = openai.AsyncOpenAI(
- **client_params
- ).completions
- else:
- values["openai_api_base"] = values["octoai_api_base"]
- values["openai_api_key"] = values["octoai_api_token"].get_secret_value()
- values["client"] = openai.Completion
- except ImportError:
- raise ImportError(
- "Could not import openai python package. "
- "Please install it with `pip install openai`."
- )
-
- if "endpoint_url" in values["model_kwargs"]:
- raise ValueError(
- "`endpoint_url` was deprecated, please use `octoai_api_base`."
- )
-
- return values
diff --git a/libs/community/langchain_community/llms/ollama.py b/libs/community/langchain_community/llms/ollama.py
deleted file mode 100644
index e6584ae218..0000000000
--- a/libs/community/langchain_community/llms/ollama.py
+++ /dev/null
@@ -1,512 +0,0 @@
-from __future__ import annotations
-
-import json
-from typing import (
- Any,
- AsyncIterator,
- Callable,
- Dict,
- Iterator,
- List,
- Mapping,
- Optional,
- Tuple,
- Union,
-)
-
-import aiohttp
-import requests
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models import BaseLanguageModel
-from langchain_core.language_models.llms import BaseLLM
-from langchain_core.outputs import GenerationChunk, LLMResult
-from pydantic import ConfigDict
-
-
-def _stream_response_to_generation_chunk(
- stream_response: str,
-) -> GenerationChunk:
- """Convert a stream response to a generation chunk."""
- parsed_response = json.loads(stream_response)
- generation_info = parsed_response if parsed_response.get("done") is True else None
- return GenerationChunk(
- text=parsed_response.get("response", ""), generation_info=generation_info
- )
-
-
-class OllamaEndpointNotFoundError(Exception):
- """Raised when the Ollama endpoint is not found."""
-
-
-class _OllamaCommon(BaseLanguageModel):
- base_url: str = "http://localhost:11434"
- """Base url the model is hosted under."""
-
- model: str = "llama2"
- """Model name to use."""
-
- mirostat: Optional[int] = None
- """Enable Mirostat sampling for controlling perplexity.
- (default: 0, 0 = disabled, 1 = Mirostat, 2 = Mirostat 2.0)"""
-
- mirostat_eta: Optional[float] = None
- """Influences how quickly the algorithm responds to feedback
- from the generated text. A lower learning rate will result in
- slower adjustments, while a higher learning rate will make
- the algorithm more responsive. (Default: 0.1)"""
-
- mirostat_tau: Optional[float] = None
- """Controls the balance between coherence and diversity
- of the output. A lower value will result in more focused and
- coherent text. (Default: 5.0)"""
-
- num_ctx: Optional[int] = None
- """Sets the size of the context window used to generate the
- next token. (Default: 2048) """
-
- num_gpu: Optional[int] = None
- """The number of GPUs to use. On macOS it defaults to 1 to
- enable metal support, 0 to disable."""
-
- num_thread: Optional[int] = None
- """Sets the number of threads to use during computation.
- By default, Ollama will detect this for optimal performance.
- It is recommended to set this value to the number of physical
- CPU cores your system has (as opposed to the logical number of cores)."""
-
- num_predict: Optional[int] = None
- """Maximum number of tokens to predict when generating text.
- (Default: 128, -1 = infinite generation, -2 = fill context)"""
-
- repeat_last_n: Optional[int] = None
- """Sets how far back for the model to look back to prevent
- repetition. (Default: 64, 0 = disabled, -1 = num_ctx)"""
-
- repeat_penalty: Optional[float] = None
- """Sets how strongly to penalize repetitions. A higher value (e.g., 1.5)
- will penalize repetitions more strongly, while a lower value (e.g., 0.9)
- will be more lenient. (Default: 1.1)"""
-
- temperature: Optional[float] = None
- """The temperature of the model. Increasing the temperature will
- make the model answer more creatively. (Default: 0.8)"""
-
- stop: Optional[List[str]] = None
- """Sets the stop tokens to use."""
-
- tfs_z: Optional[float] = None
- """Tail free sampling is used to reduce the impact of less probable
- tokens from the output. A higher value (e.g., 2.0) will reduce the
- impact more, while a value of 1.0 disables this setting. (default: 1)"""
-
- top_k: Optional[int] = None
- """Reduces the probability of generating nonsense. A higher value (e.g. 100)
- will give more diverse answers, while a lower value (e.g. 10)
- will be more conservative. (Default: 40)"""
-
- top_p: Optional[float] = None
- """Works together with top-k. A higher value (e.g., 0.95) will lead
- to more diverse text, while a lower value (e.g., 0.5) will
- generate more focused and conservative text. (Default: 0.9)"""
-
- system: Optional[str] = None
- """system prompt (overrides what is defined in the Modelfile)"""
-
- template: Optional[str] = None
- """full prompt or prompt template (overrides what is defined in the Modelfile)"""
-
- format: Optional[str] = None
- """Specify the format of the output (e.g., json)"""
-
- timeout: Optional[int] = None
- """Timeout for the request stream"""
-
- keep_alive: Optional[Union[int, str]] = None
- """How long the model will stay loaded into memory.
-
- The parameter (Default: 5 minutes) can be set to:
- 1. a duration string in Golang (such as "10m" or "24h");
- 2. a number in seconds (such as 3600);
- 3. any negative number which will keep the model loaded \
- in memory (e.g. -1 or "-1m");
- 4. 0 which will unload the model immediately after generating a response;
-
- See the [Ollama documents](https://github.com/ollama/ollama/blob/main/docs/faq.md#how-do-i-keep-a-model-loaded-in-memory-or-make-it-unload-immediately)"""
-
- raw: Optional[bool] = None
- """raw or not."""
-
- headers: Optional[dict] = None
- """Additional headers to pass to endpoint (e.g. Authorization, Referer).
- This is useful when Ollama is hosted on cloud services that require
- tokens for authentication.
- """
-
- auth: Union[Callable, Tuple, None] = None
- """Additional auth tuple or callable to enable Basic/Digest/Custom HTTP Auth.
- Expects the same format, type and values as requests.request auth parameter."""
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling Ollama."""
- return {
- "model": self.model,
- "format": self.format,
- "options": {
- "mirostat": self.mirostat,
- "mirostat_eta": self.mirostat_eta,
- "mirostat_tau": self.mirostat_tau,
- "num_ctx": self.num_ctx,
- "num_gpu": self.num_gpu,
- "num_thread": self.num_thread,
- "num_predict": self.num_predict,
- "repeat_last_n": self.repeat_last_n,
- "repeat_penalty": self.repeat_penalty,
- "temperature": self.temperature,
- "stop": self.stop,
- "tfs_z": self.tfs_z,
- "top_k": self.top_k,
- "top_p": self.top_p,
- },
- "system": self.system,
- "template": self.template,
- "keep_alive": self.keep_alive,
- "raw": self.raw,
- }
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {**{"model": self.model, "format": self.format}, **self._default_params}
-
- def _create_generate_stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- images: Optional[List[str]] = None,
- **kwargs: Any,
- ) -> Iterator[str]:
- payload = {"prompt": prompt, "images": images}
- yield from self._create_stream(
- payload=payload,
- stop=stop,
- api_url=f"{self.base_url}/api/generate",
- **kwargs,
- )
-
- async def _acreate_generate_stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- images: Optional[List[str]] = None,
- **kwargs: Any,
- ) -> AsyncIterator[str]:
- payload = {"prompt": prompt, "images": images}
- async for item in self._acreate_stream(
- payload=payload,
- stop=stop,
- api_url=f"{self.base_url}/api/generate",
- **kwargs,
- ):
- yield item
-
- def _create_stream(
- self,
- api_url: str,
- payload: Any,
- stop: Optional[List[str]] = None,
- **kwargs: Any,
- ) -> Iterator[str]:
- if self.stop is not None and stop is not None:
- raise ValueError("`stop` found in both the input and default params.")
- elif self.stop is not None:
- stop = self.stop
-
- params = self._default_params
-
- for key in self._default_params:
- if key in kwargs:
- params[key] = kwargs[key]
-
- if "options" in kwargs:
- params["options"] = kwargs["options"]
- else:
- params["options"] = {
- **params["options"],
- "stop": stop,
- **{k: v for k, v in kwargs.items() if k not in self._default_params},
- }
-
- if payload.get("messages"):
- request_payload = {"messages": payload.get("messages", []), **params}
- else:
- request_payload = {
- "prompt": payload.get("prompt"),
- "images": payload.get("images", []),
- **params,
- }
- response = requests.post(
- url=api_url,
- headers={
- "Content-Type": "application/json",
- **(self.headers if isinstance(self.headers, dict) else {}),
- },
- auth=self.auth,
- json=request_payload,
- stream=True,
- timeout=self.timeout,
- )
- response.encoding = "utf-8"
- if response.status_code != 200:
- if response.status_code == 404:
- raise OllamaEndpointNotFoundError(
- "Ollama call failed with status code 404. "
- "Maybe your model is not found "
- f"and you should pull the model with `ollama pull {self.model}`."
- )
- else:
- optional_detail = response.text
- raise ValueError(
- f"Ollama call failed with status code {response.status_code}."
- f" Details: {optional_detail}"
- )
- return response.iter_lines(decode_unicode=True)
-
- async def _acreate_stream(
- self,
- api_url: str,
- payload: Any,
- stop: Optional[List[str]] = None,
- **kwargs: Any,
- ) -> AsyncIterator[str]:
- if self.stop is not None and stop is not None:
- raise ValueError("`stop` found in both the input and default params.")
- elif self.stop is not None:
- stop = self.stop
-
- params = self._default_params
-
- for key in self._default_params:
- if key in kwargs:
- params[key] = kwargs[key]
-
- if "options" in kwargs:
- params["options"] = kwargs["options"]
- else:
- params["options"] = {
- **params["options"],
- "stop": stop,
- **{k: v for k, v in kwargs.items() if k not in self._default_params},
- }
-
- if payload.get("messages"):
- request_payload = {"messages": payload.get("messages", []), **params}
- else:
- request_payload = {
- "prompt": payload.get("prompt"),
- "images": payload.get("images", []),
- **params,
- }
-
- async with aiohttp.ClientSession() as session:
- async with session.post(
- url=api_url,
- headers={
- "Content-Type": "application/json",
- **(self.headers if isinstance(self.headers, dict) else {}),
- },
- auth=self.auth, # type: ignore[arg-type,unused-ignore]
- json=request_payload,
- timeout=self.timeout, # type: ignore[arg-type,unused-ignore]
- ) as response:
- if response.status != 200:
- if response.status == 404:
- raise OllamaEndpointNotFoundError(
- "Ollama call failed with status code 404."
- )
- else:
- optional_detail = response.text
- raise ValueError(
- f"Ollama call failed with status code {response.status}."
- f" Details: {optional_detail}"
- )
- async for line in response.content:
- yield line.decode("utf-8")
-
- def _stream_with_aggregation(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- verbose: bool = False,
- **kwargs: Any,
- ) -> GenerationChunk:
- final_chunk: Optional[GenerationChunk] = None
- for stream_resp in self._create_generate_stream(prompt, stop, **kwargs):
- if stream_resp:
- chunk = _stream_response_to_generation_chunk(stream_resp)
- if final_chunk is None:
- final_chunk = chunk
- else:
- final_chunk += chunk
- if run_manager:
- run_manager.on_llm_new_token(
- chunk.text,
- verbose=verbose,
- )
- if final_chunk is None:
- raise ValueError("No data received from Ollama stream.")
-
- return final_chunk
-
- async def _astream_with_aggregation(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- verbose: bool = False,
- **kwargs: Any,
- ) -> GenerationChunk:
- final_chunk: Optional[GenerationChunk] = None
- async for stream_resp in self._acreate_generate_stream(prompt, stop, **kwargs):
- if stream_resp:
- chunk = _stream_response_to_generation_chunk(stream_resp)
- if final_chunk is None:
- final_chunk = chunk
- else:
- final_chunk += chunk
- if run_manager:
- await run_manager.on_llm_new_token(
- chunk.text,
- verbose=verbose,
- )
- if final_chunk is None:
- raise ValueError("No data received from Ollama stream.")
-
- return final_chunk
-
-
-@deprecated(
- since="0.3.1",
- removal="1.0.0",
- alternative_import="langchain_ollama.OllamaLLM",
-)
-class Ollama(BaseLLM, _OllamaCommon):
- """Ollama locally runs large language models.
- To use, follow the instructions at https://ollama.ai/.
- Example:
- .. code-block:: python
- from langchain_community.llms import Ollama
- ollama = Ollama(model="llama2")
- """
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "ollama-llm"
-
- def _generate( # type: ignore[override]
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- images: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call out to Ollama's generate endpoint.
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
- Returns:
- The string generated by the model.
- Example:
- .. code-block:: python
- response = ollama("Tell me a joke.")
- """
- # TODO: add caching here.
- generations = []
- for prompt in prompts:
- final_chunk = super()._stream_with_aggregation(
- prompt,
- stop=stop,
- images=images,
- run_manager=run_manager,
- verbose=self.verbose,
- **kwargs,
- )
- generations.append([final_chunk])
- return LLMResult(generations=generations) # type: ignore[arg-type]
-
- async def _agenerate( # type: ignore[override]
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- images: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call out to Ollama's generate endpoint.
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
- Returns:
- The string generated by the model.
- Example:
- .. code-block:: python
- response = ollama("Tell me a joke.")
- """
- # TODO: add caching here.
- generations = []
- for prompt in prompts:
- final_chunk = await super()._astream_with_aggregation(
- prompt,
- stop=stop,
- images=images,
- run_manager=run_manager, # type: ignore[arg-type]
- verbose=self.verbose,
- **kwargs,
- )
- generations.append([final_chunk])
- return LLMResult(generations=generations) # type: ignore[arg-type]
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- for stream_resp in self._create_generate_stream(prompt, stop, **kwargs):
- if stream_resp:
- chunk = _stream_response_to_generation_chunk(stream_resp)
- if run_manager:
- run_manager.on_llm_new_token(
- chunk.text,
- verbose=self.verbose,
- )
- yield chunk
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- async for stream_resp in self._acreate_generate_stream(prompt, stop, **kwargs):
- if stream_resp:
- chunk = _stream_response_to_generation_chunk(stream_resp)
- if run_manager:
- await run_manager.on_llm_new_token(
- chunk.text,
- verbose=self.verbose,
- )
- yield chunk
diff --git a/libs/community/langchain_community/llms/opaqueprompts.py b/libs/community/langchain_community/llms/opaqueprompts.py
deleted file mode 100644
index 46a2a2b36f..0000000000
--- a/libs/community/langchain_community/llms/opaqueprompts.py
+++ /dev/null
@@ -1,117 +0,0 @@
-import logging
-from typing import Any, Dict, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models import BaseLanguageModel
-from langchain_core.language_models.llms import LLM
-from langchain_core.messages import AIMessage
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import ConfigDict
-
-logger = logging.getLogger(__name__)
-
-
-class OpaquePrompts(LLM):
- """LLM that uses OpaquePrompts to sanitize prompts.
-
- Wraps another LLM and sanitizes prompts before passing it to the LLM, then
- de-sanitizes the response.
-
- To use, you should have the ``opaqueprompts`` python package installed,
- and the environment variable ``OPAQUEPROMPTS_API_KEY`` set with
- your API key, or pass it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import OpaquePrompts
- from langchain_community.chat_models import ChatOpenAI
-
- op_llm = OpaquePrompts(base_llm=ChatOpenAI())
- """
-
- base_llm: BaseLanguageModel
- """The base LLM to use."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validates that the OpaquePrompts API key and the Python package exist."""
- try:
- import opaqueprompts as op
- except ImportError:
- raise ImportError(
- "Could not import the `opaqueprompts` Python package, "
- "please install it with `pip install opaqueprompts`."
- )
- if op.__package__ is None:
- raise ValueError(
- "Could not properly import `opaqueprompts`, "
- "opaqueprompts.__package__ is None."
- )
-
- api_key = get_from_dict_or_env(
- values, "opaqueprompts_api_key", "OPAQUEPROMPTS_API_KEY", default=""
- )
- if not api_key:
- raise ValueError(
- "Could not find OPAQUEPROMPTS_API_KEY in the environment. "
- "Please set it to your OpaquePrompts API key."
- "You can get it by creating an account on the OpaquePrompts website: "
- "https://opaqueprompts.opaque.co/ ."
- )
- return values
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call base LLM with sanitization before and de-sanitization after.
-
- Args:
- prompt: The prompt to pass into the model.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = op_llm.invoke("Tell me a joke.")
- """
- import opaqueprompts as op
-
- _run_manager = run_manager or CallbackManagerForLLMRun.get_noop_manager()
-
- # sanitize the prompt by replacing the sensitive information with a placeholder
- sanitize_response: op.SanitizeResponse = op.sanitize([prompt])
- sanitized_prompt_value_str = sanitize_response.sanitized_texts[0]
-
- # TODO: Add in callbacks once child runs for LLMs are supported by LangSmith.
- # call the LLM with the sanitized prompt and get the response
- llm_response = self.base_llm.bind(stop=stop).invoke(
- sanitized_prompt_value_str,
- )
- if isinstance(llm_response, AIMessage):
- llm_response = llm_response.content
-
- # desanitize the response by restoring the original sensitive information
- desanitize_response: op.DesanitizeResponse = op.desanitize(
- llm_response,
- secure_context=sanitize_response.secure_context,
- )
- return desanitize_response.desanitized_text
-
- @property
- def _llm_type(self) -> str:
- """Return type of LLM.
-
- This is an override of the base class method.
- """
- return "opaqueprompts"
diff --git a/libs/community/langchain_community/llms/openai.py b/libs/community/langchain_community/llms/openai.py
deleted file mode 100644
index 525599945f..0000000000
--- a/libs/community/langchain_community/llms/openai.py
+++ /dev/null
@@ -1,1258 +0,0 @@
-from __future__ import annotations
-
-import logging
-import os
-import sys
-import warnings
-from typing import (
- AbstractSet,
- Any,
- AsyncIterator,
- Awaitable,
- Callable,
- Collection,
- Dict,
- Iterator,
- List,
- Literal,
- Mapping,
- Optional,
- Set,
- Tuple,
- Union,
-)
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import BaseLLM, create_base_retry_decorator
-from langchain_core.outputs import Generation, GenerationChunk, LLMResult
-from langchain_core.utils import (
- get_from_dict_or_env,
- get_pydantic_field_names,
- pre_init,
-)
-from langchain_core.utils.pydantic import get_fields
-from langchain_core.utils.utils import _build_model_kwargs
-from pydantic import ConfigDict, Field, model_validator
-
-from langchain_community.utils.openai import is_openai_v1
-
-logger = logging.getLogger(__name__)
-
-
-def update_token_usage(
- keys: Set[str], response: Dict[str, Any], token_usage: Dict[str, Any]
-) -> None:
- """Update token usage."""
- _keys_to_use = keys.intersection(response["usage"])
- for _key in _keys_to_use:
- if _key not in token_usage:
- token_usage[_key] = response["usage"][_key]
- else:
- token_usage[_key] += response["usage"][_key]
-
-
-def _stream_response_to_generation_chunk(
- stream_response: Dict[str, Any],
-) -> GenerationChunk:
- """Convert a stream response to a generation chunk."""
- if not stream_response["choices"]:
- return GenerationChunk(text="")
- return GenerationChunk(
- text=stream_response["choices"][0]["text"],
- generation_info=dict(
- finish_reason=stream_response["choices"][0].get("finish_reason", None),
- logprobs=stream_response["choices"][0].get("logprobs", None),
- ),
- )
-
-
-def _update_response(response: Dict[str, Any], stream_response: Dict[str, Any]) -> None:
- """Update response from the stream response."""
- response["choices"][0]["text"] += stream_response["choices"][0]["text"]
- response["choices"][0]["finish_reason"] = stream_response["choices"][0].get(
- "finish_reason", None
- )
- response["choices"][0]["logprobs"] = stream_response["choices"][0]["logprobs"]
-
-
-def _streaming_response_template() -> Dict[str, Any]:
- return {
- "choices": [
- {
- "text": "",
- "finish_reason": None,
- "logprobs": None,
- }
- ]
- }
-
-
-def _create_retry_decorator(
- llm: Union[BaseOpenAI, OpenAIChat],
- run_manager: Optional[
- Union[AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun]
- ] = None,
-) -> Callable[[Any], Any]:
- import openai
-
- errors = [
- openai.error.Timeout,
- openai.error.APIError,
- openai.error.APIConnectionError,
- openai.error.RateLimitError,
- openai.error.ServiceUnavailableError,
- ]
- return create_base_retry_decorator(
- error_types=errors, max_retries=llm.max_retries, run_manager=run_manager
- )
-
-
-def completion_with_retry(
- llm: Union[BaseOpenAI, OpenAIChat],
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
-) -> Any:
- """Use tenacity to retry the completion call."""
- if is_openai_v1():
- return llm.client.create(**kwargs)
-
- retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)
-
- @retry_decorator
- def _completion_with_retry(**kwargs: Any) -> Any:
- return llm.client.create(**kwargs)
-
- return _completion_with_retry(**kwargs)
-
-
-async def acompletion_with_retry(
- llm: Union[BaseOpenAI, OpenAIChat],
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
-) -> Any:
- """Use tenacity to retry the async completion call."""
- if is_openai_v1():
- return await llm.async_client.create(**kwargs)
-
- retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)
-
- @retry_decorator
- async def _completion_with_retry(**kwargs: Any) -> Any:
- # Use OpenAI's async api https://github.com/openai/openai-python#async-api
- return await llm.client.acreate(**kwargs)
-
- return await _completion_with_retry(**kwargs)
-
-
-class BaseOpenAI(BaseLLM):
- """Base OpenAI large language model class."""
-
- @property
- def lc_secrets(self) -> Dict[str, str]:
- return {"openai_api_key": "OPENAI_API_KEY"}
-
- @classmethod
- def get_lc_namespace(cls) -> List[str]:
- """Get the namespace of the langchain object."""
- return ["langchain", "llms", "openai"]
-
- @property
- def lc_attributes(self) -> Dict[str, Any]:
- attributes: Dict[str, Any] = {}
- if self.openai_api_base:
- attributes["openai_api_base"] = self.openai_api_base
-
- if self.openai_organization:
- attributes["openai_organization"] = self.openai_organization
-
- if self.openai_proxy:
- attributes["openai_proxy"] = self.openai_proxy
-
- return attributes
-
- @classmethod
- def is_lc_serializable(cls) -> bool:
- return True
-
- client: Any = Field(default=None, exclude=True) #: :meta private:
- async_client: Any = Field(default=None, exclude=True) #: :meta private:
- model_name: str = Field(default="gpt-3.5-turbo-instruct", alias="model")
- """Model name to use."""
- temperature: float = 0.7
- """What sampling temperature to use."""
- max_tokens: int = 256
- """The maximum number of tokens to generate in the completion.
- -1 returns as many tokens as possible given the prompt and
- the models maximal context size."""
- top_p: float = 1
- """Total probability mass of tokens to consider at each step."""
- frequency_penalty: float = 0
- """Penalizes repeated tokens according to frequency."""
- presence_penalty: float = 0
- """Penalizes repeated tokens."""
- n: int = 1
- """How many completions to generate for each prompt."""
- best_of: int = 1
- """Generates best_of completions server-side and returns the "best"."""
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `create` call not explicitly specified."""
- # When updating this to use a SecretStr
- # Check for classes that derive from this class (as some of them
- # may assume openai_api_key is a str)
- openai_api_key: Optional[str] = Field(default=None, alias="api_key")
- """Automatically inferred from env var `OPENAI_API_KEY` if not provided."""
- openai_api_base: Optional[str] = Field(default=None, alias="base_url")
- """Base URL path for API requests, leave blank if not using a proxy or service
- emulator."""
- openai_organization: Optional[str] = Field(default=None, alias="organization")
- """Automatically inferred from env var `OPENAI_ORG_ID` if not provided."""
- # to support explicit proxy for OpenAI
- openai_proxy: Optional[str] = None
- batch_size: int = 20
- """Batch size to use when passing multiple documents to generate."""
- request_timeout: Union[float, Tuple[float, float], Any, None] = Field(
- default=None, alias="timeout"
- )
- """Timeout for requests to OpenAI completion API. Can be float, httpx.Timeout or
- None."""
- logit_bias: Optional[Dict[str, float]] = Field(default_factory=dict) # type: ignore[arg-type]
- """Adjust the probability of specific tokens being generated."""
- max_retries: int = 2
- """Maximum number of retries to make when generating."""
- streaming: bool = False
- """Whether to stream the results or not."""
- allowed_special: Union[Literal["all"], AbstractSet[str]] = set()
- """Set of special tokens that are allowed。"""
- disallowed_special: Union[Literal["all"], Collection[str]] = "all"
- """Set of special tokens that are not allowed。"""
- tiktoken_model_name: Optional[str] = None
- """The model name to pass to tiktoken when using this class.
- Tiktoken is used to count the number of tokens in documents to constrain
- them to be under a certain limit. By default, when set to None, this will
- be the same as the embedding model name. However, there are some cases
- where you may want to use this Embedding class with a model name not
- supported by tiktoken. This can include when using Azure embeddings or
- when using one of the many model providers that expose an OpenAI-like
- API but with different models. In those cases, in order to avoid erroring
- when tiktoken is called, you can specify a model name to use here."""
- default_headers: Union[Mapping[str, str], None] = None
- default_query: Union[Mapping[str, object], None] = None
- # Configure a custom httpx client. See the
- # [httpx documentation](https://www.python-httpx.org/api/#client) for more details.
- http_client: Union[Any, None] = None
- """Optional httpx.Client."""
-
- def __new__(cls, **data: Any) -> Union[OpenAIChat, BaseOpenAI]: # type: ignore[misc]
- """Initialize the OpenAI object."""
- model_name = data.get("model_name", "")
- if (
- model_name.startswith("gpt-3.5-turbo") or model_name.startswith("gpt-4")
- ) and "-instruct" not in model_name:
- warnings.warn(
- "You are trying to use a chat model. This way of initializing it is "
- "no longer supported. Instead, please use: "
- "`from langchain_community.chat_models import ChatOpenAI`"
- )
- return OpenAIChat(**data)
- return super().__new__(cls)
-
- model_config = ConfigDict(
- populate_by_name=True,
- )
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = get_pydantic_field_names(cls)
- values = _build_model_kwargs(values, all_required_field_names)
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- if values["n"] < 1:
- raise ValueError("n must be at least 1.")
- if values["streaming"] and values["n"] > 1:
- raise ValueError("Cannot stream results when n > 1.")
- if values["streaming"] and values["best_of"] > 1:
- raise ValueError("Cannot stream results when best_of > 1.")
-
- values["openai_api_key"] = get_from_dict_or_env(
- values, "openai_api_key", "OPENAI_API_KEY"
- )
- values["openai_api_base"] = values["openai_api_base"] or os.getenv(
- "OPENAI_API_BASE"
- )
- values["openai_proxy"] = get_from_dict_or_env(
- values,
- "openai_proxy",
- "OPENAI_PROXY",
- default="",
- )
- values["openai_organization"] = (
- values["openai_organization"]
- or os.getenv("OPENAI_ORG_ID")
- or os.getenv("OPENAI_ORGANIZATION")
- )
- try:
- import openai
- except ImportError:
- raise ImportError(
- "Could not import openai python package. "
- "Please install it with `pip install openai`."
- )
-
- if is_openai_v1():
- client_params = {
- "api_key": values["openai_api_key"],
- "organization": values["openai_organization"],
- "base_url": values["openai_api_base"],
- "timeout": values["request_timeout"],
- "max_retries": values["max_retries"],
- "default_headers": values["default_headers"],
- "default_query": values["default_query"],
- "http_client": values["http_client"],
- }
- if not values.get("client"):
- values["client"] = openai.OpenAI(**client_params).completions
- if not values.get("async_client"):
- values["async_client"] = openai.AsyncOpenAI(**client_params).completions
- elif not values.get("client"):
- values["client"] = openai.Completion
- else:
- pass
-
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling OpenAI API."""
- normal_params: Dict[str, Any] = {
- "temperature": self.temperature,
- "top_p": self.top_p,
- "frequency_penalty": self.frequency_penalty,
- "presence_penalty": self.presence_penalty,
- "n": self.n,
- "logit_bias": self.logit_bias,
- }
-
- if self.max_tokens is not None:
- normal_params["max_tokens"] = self.max_tokens
- if self.request_timeout is not None and not is_openai_v1():
- normal_params["request_timeout"] = self.request_timeout
-
- # Azure gpt-35-turbo doesn't support best_of
- # don't specify best_of if it is 1
- if self.best_of > 1:
- normal_params["best_of"] = self.best_of
-
- return {**normal_params, **self.model_kwargs}
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- params = {**self._invocation_params, **kwargs, "stream": True}
- self.get_sub_prompts(params, [prompt], stop) # this mutates params
- for stream_resp in completion_with_retry(
- self, prompt=prompt, run_manager=run_manager, **params
- ):
- if not isinstance(stream_resp, dict):
- stream_resp = stream_resp.dict()
- chunk = _stream_response_to_generation_chunk(stream_resp)
- if run_manager:
- run_manager.on_llm_new_token(
- chunk.text,
- chunk=chunk,
- verbose=self.verbose,
- logprobs=chunk.generation_info["logprobs"]
- if chunk.generation_info
- else None,
- )
- yield chunk
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- params = {**self._invocation_params, **kwargs, "stream": True}
- self.get_sub_prompts(params, [prompt], stop) # this mutates params
- async for stream_resp in await acompletion_with_retry(
- self, prompt=prompt, run_manager=run_manager, **params
- ):
- if not isinstance(stream_resp, dict):
- stream_resp = stream_resp.dict()
- chunk = _stream_response_to_generation_chunk(stream_resp)
- if run_manager:
- await run_manager.on_llm_new_token(
- chunk.text,
- chunk=chunk,
- verbose=self.verbose,
- logprobs=chunk.generation_info["logprobs"]
- if chunk.generation_info
- else None,
- )
- yield chunk
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call out to OpenAI's endpoint with k unique prompts.
-
- Args:
- prompts: The prompts to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The full LLM output.
-
- Example:
- .. code-block:: python
-
- response = openai.generate(["Tell me a joke."])
- """
- # TODO: write a unit test for this
- params = self._invocation_params
- params = {**params, **kwargs}
- sub_prompts = self.get_sub_prompts(params, prompts, stop)
- choices = []
- token_usage: Dict[str, int] = {}
- # Get the token usage from the response.
- # Includes prompt, completion, and total tokens used.
- _keys = {"completion_tokens", "prompt_tokens", "total_tokens"}
- system_fingerprint: Optional[str] = None
- for _prompts in sub_prompts:
- if self.streaming:
- if len(_prompts) > 1:
- raise ValueError("Cannot stream results with multiple prompts.")
-
- generation: Optional[GenerationChunk] = None
- for chunk in self._stream(_prompts[0], stop, run_manager, **kwargs):
- if generation is None:
- generation = chunk
- else:
- generation += chunk
- assert generation is not None
- choices.append(
- {
- "text": generation.text,
- "finish_reason": generation.generation_info.get("finish_reason")
- if generation.generation_info
- else None,
- "logprobs": generation.generation_info.get("logprobs")
- if generation.generation_info
- else None,
- }
- )
- else:
- response = completion_with_retry(
- self, prompt=_prompts, run_manager=run_manager, **params
- )
- if not isinstance(response, dict):
- # V1 client returns the response in an PyDantic object instead of
- # dict. For the transition period, we deep convert it to dict.
- response = response.dict()
-
- choices.extend(response["choices"])
- update_token_usage(_keys, response, token_usage)
- if not system_fingerprint:
- system_fingerprint = response.get("system_fingerprint")
- return self.create_llm_result(
- choices,
- prompts,
- params,
- token_usage,
- system_fingerprint=system_fingerprint,
- )
-
- async def _agenerate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call out to OpenAI's endpoint async with k unique prompts."""
- params = self._invocation_params
- params = {**params, **kwargs}
- sub_prompts = self.get_sub_prompts(params, prompts, stop)
- choices = []
- token_usage: Dict[str, int] = {}
- # Get the token usage from the response.
- # Includes prompt, completion, and total tokens used.
- _keys = {"completion_tokens", "prompt_tokens", "total_tokens"}
- system_fingerprint: Optional[str] = None
- for _prompts in sub_prompts:
- if self.streaming:
- if len(_prompts) > 1:
- raise ValueError("Cannot stream results with multiple prompts.")
-
- generation: Optional[GenerationChunk] = None
- async for chunk in self._astream(
- _prompts[0], stop, run_manager, **kwargs
- ):
- if generation is None:
- generation = chunk
- else:
- generation += chunk
- assert generation is not None
- choices.append(
- {
- "text": generation.text,
- "finish_reason": generation.generation_info.get("finish_reason")
- if generation.generation_info
- else None,
- "logprobs": generation.generation_info.get("logprobs")
- if generation.generation_info
- else None,
- }
- )
- else:
- response = await acompletion_with_retry(
- self, prompt=_prompts, run_manager=run_manager, **params
- )
- if not isinstance(response, dict):
- response = response.dict()
- choices.extend(response["choices"])
- update_token_usage(_keys, response, token_usage)
- return self.create_llm_result(
- choices,
- prompts,
- params,
- token_usage,
- system_fingerprint=system_fingerprint,
- )
-
- def get_sub_prompts(
- self,
- params: Dict[str, Any],
- prompts: List[str],
- stop: Optional[List[str]] = None,
- ) -> List[List[str]]:
- """Get the sub prompts for llm call."""
- if stop is not None:
- if "stop" in params:
- raise ValueError("`stop` found in both the input and default params.")
- params["stop"] = stop
- if params["max_tokens"] == -1:
- if len(prompts) != 1:
- raise ValueError(
- "max_tokens set to -1 not supported for multiple inputs."
- )
- params["max_tokens"] = self.max_tokens_for_prompt(prompts[0])
- sub_prompts = [
- prompts[i : i + self.batch_size]
- for i in range(0, len(prompts), self.batch_size)
- ]
- return sub_prompts
-
- def create_llm_result(
- self,
- choices: Any,
- prompts: List[str],
- params: Dict[str, Any],
- token_usage: Dict[str, int],
- *,
- system_fingerprint: Optional[str] = None,
- ) -> LLMResult:
- """Create the LLMResult from the choices and prompts."""
- generations = []
- n = params.get("n", self.n)
- for i, _ in enumerate(prompts):
- sub_choices = choices[i * n : (i + 1) * n]
- generations.append(
- [
- Generation(
- text=choice["text"],
- generation_info=dict(
- finish_reason=choice.get("finish_reason"),
- logprobs=choice.get("logprobs"),
- ),
- )
- for choice in sub_choices
- ]
- )
- llm_output = {"token_usage": token_usage, "model_name": self.model_name}
- if system_fingerprint:
- llm_output["system_fingerprint"] = system_fingerprint
- return LLMResult(generations=generations, llm_output=llm_output)
-
- @property
- def _invocation_params(self) -> Dict[str, Any]:
- """Get the parameters used to invoke the model."""
- openai_creds: Dict[str, Any] = {}
- if not is_openai_v1():
- openai_creds.update(
- {
- "api_key": self.openai_api_key,
- "api_base": self.openai_api_base,
- "organization": self.openai_organization,
- }
- )
- if self.openai_proxy:
- import openai
-
- openai.proxy = {"http": self.openai_proxy, "https": self.openai_proxy}
- return {**openai_creds, **self._default_params}
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {**{"model_name": self.model_name}, **self._default_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "openai"
-
- def get_token_ids(self, text: str) -> List[int]:
- """Get the token IDs using the tiktoken package."""
- # tiktoken NOT supported for Python < 3.8
- if sys.version_info[1] < 8:
- return super().get_num_tokens(text)
- try:
- import tiktoken
- except ImportError:
- raise ImportError(
- "Could not import tiktoken python package. "
- "This is needed in order to calculate get_num_tokens. "
- "Please install it with `pip install tiktoken`."
- )
-
- model_name = self.tiktoken_model_name or self.model_name
- try:
- enc = tiktoken.encoding_for_model(model_name)
- except KeyError:
- logger.warning("Warning: model not found. Using cl100k_base encoding.")
- model = "cl100k_base"
- enc = tiktoken.get_encoding(model)
-
- return enc.encode(
- text,
- allowed_special=self.allowed_special,
- disallowed_special=self.disallowed_special,
- )
-
- @staticmethod
- def modelname_to_contextsize(modelname: str) -> int:
- """Calculate the maximum number of tokens possible to generate for a model.
-
- Args:
- modelname: The modelname we want to know the context size for.
-
- Returns:
- The maximum context size
-
- Example:
- .. code-block:: python
-
- max_tokens = openai.modelname_to_contextsize("gpt-3.5-turbo-instruct")
- """
- model_token_mapping = {
- "gpt-4o": 128_000,
- "gpt-4o-2024-05-13": 128_000,
- "gpt-4": 8192,
- "gpt-4-0314": 8192,
- "gpt-4-0613": 8192,
- "gpt-4-32k": 32768,
- "gpt-4-32k-0314": 32768,
- "gpt-4-32k-0613": 32768,
- "gpt-3.5-turbo": 4096,
- "gpt-3.5-turbo-0301": 4096,
- "gpt-3.5-turbo-0613": 4096,
- "gpt-3.5-turbo-16k": 16385,
- "gpt-3.5-turbo-16k-0613": 16385,
- "gpt-3.5-turbo-instruct": 4096,
- "text-ada-001": 2049,
- "ada": 2049,
- "text-babbage-001": 2040,
- "babbage": 2049,
- "text-curie-001": 2049,
- "curie": 2049,
- "davinci": 2049,
- "text-davinci-003": 4097,
- "text-davinci-002": 4097,
- "code-davinci-002": 8001,
- "code-davinci-001": 8001,
- "code-cushman-002": 2048,
- "code-cushman-001": 2048,
- }
-
- # handling finetuned models
- if "ft-" in modelname:
- modelname = modelname.split(":")[0]
-
- context_size = model_token_mapping.get(modelname, None)
-
- if context_size is None:
- raise ValueError(
- f"Unknown model: {modelname}. Please provide a valid OpenAI model name."
- "Known models are: " + ", ".join(model_token_mapping.keys())
- )
-
- return context_size
-
- @property
- def max_context_size(self) -> int:
- """Get max context size for this model."""
- return self.modelname_to_contextsize(self.model_name)
-
- def max_tokens_for_prompt(self, prompt: str) -> int:
- """Calculate the maximum number of tokens possible to generate for a prompt.
-
- Args:
- prompt: The prompt to pass into the model.
-
- Returns:
- The maximum number of tokens to generate for a prompt.
-
- Example:
- .. code-block:: python
-
- max_tokens = openai.max_tokens_for_prompt("Tell me a joke.")
- """
- num_tokens = self.get_num_tokens(prompt)
- return self.max_context_size - num_tokens
-
-
-@deprecated(since="0.0.10", removal="1.0", alternative_import="langchain_openai.OpenAI")
-class OpenAI(BaseOpenAI):
- """OpenAI large language models.
-
- To use, you should have the ``openai`` python package installed, and the
- environment variable ``OPENAI_API_KEY`` set with your API key.
-
- Any parameters that are valid to be passed to the openai.create call can be passed
- in, even if not explicitly saved on this class.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import OpenAI
- openai = OpenAI(model_name="gpt-3.5-turbo-instruct")
- """
-
- @classmethod
- def get_lc_namespace(cls) -> List[str]:
- """Get the namespace of the langchain object."""
- return ["langchain", "llms", "openai"]
-
- @property
- def _invocation_params(self) -> Dict[str, Any]:
- return {**{"model": self.model_name}, **super()._invocation_params}
-
-
-@deprecated(
- since="0.0.10", removal="1.0", alternative_import="langchain_openai.AzureOpenAI"
-)
-class AzureOpenAI(BaseOpenAI):
- """Azure-specific OpenAI large language models.
-
- To use, you should have the ``openai`` python package installed, and the
- environment variable ``OPENAI_API_KEY`` set with your API key.
-
- Any parameters that are valid to be passed to the openai.create call can be passed
- in, even if not explicitly saved on this class.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import AzureOpenAI
-
- openai = AzureOpenAI(model_name="gpt-3.5-turbo-instruct")
- """
-
- azure_endpoint: Union[str, None] = None
- """Your Azure endpoint, including the resource.
-
- Automatically inferred from env var `AZURE_OPENAI_ENDPOINT` if not provided.
-
- Example: `https://example-resource.azure.openai.com/`
- """
- deployment_name: Union[str, None] = Field(default=None, alias="azure_deployment")
- """A model deployment.
-
- If given sets the base client URL to include `/deployments/{azure_deployment}`.
- Note: this means you won't be able to use non-deployment endpoints.
- """
- openai_api_version: str = Field(default="", alias="api_version")
- """Automatically inferred from env var `OPENAI_API_VERSION` if not provided."""
- openai_api_key: Union[str, None] = Field(default=None, alias="api_key")
- """Automatically inferred from env var `AZURE_OPENAI_API_KEY` if not provided."""
- azure_ad_token: Union[str, None] = None
- """Your Azure Active Directory token.
-
- Automatically inferred from env var `AZURE_OPENAI_AD_TOKEN` if not provided.
-
- For more:
- https://www.microsoft.com/en-us/security/business/identity-access/microsoft-entra-id.
- """
- azure_ad_token_provider: Union[Callable[[], str], None] = None
- """A function that returns an Azure Active Directory token.
-
- Will be invoked on every sync request. For async requests,
- will be invoked if `azure_ad_async_token_provider` is not provided.
- """
- azure_ad_async_token_provider: Union[Callable[[], Awaitable[str]], None] = None
- """A function that returns an Azure Active Directory token.
-
- Will be invoked on every async request.
- """
- openai_api_type: str = ""
- """Legacy, for openai<1.0.0 support."""
- validate_base_url: bool = True
- """For backwards compatibility. If legacy val openai_api_base is passed in, try to
- infer if it is a base_url or azure_endpoint and update accordingly.
- """
-
- @classmethod
- def get_lc_namespace(cls) -> List[str]:
- """Get the namespace of the langchain object."""
- return ["langchain", "llms", "openai"]
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- if values["n"] < 1:
- raise ValueError("n must be at least 1.")
- if values["streaming"] and values["n"] > 1:
- raise ValueError("Cannot stream results when n > 1.")
- if values["streaming"] and values["best_of"] > 1:
- raise ValueError("Cannot stream results when best_of > 1.")
-
- # Check OPENAI_KEY for backwards compatibility.
- # TODO: Remove OPENAI_API_KEY support to avoid possible conflict when using
- # other forms of azure credentials.
- values["openai_api_key"] = (
- values["openai_api_key"]
- or os.getenv("AZURE_OPENAI_API_KEY")
- or os.getenv("OPENAI_API_KEY")
- )
-
- values["azure_endpoint"] = values["azure_endpoint"] or os.getenv(
- "AZURE_OPENAI_ENDPOINT"
- )
- values["azure_ad_token"] = values["azure_ad_token"] or os.getenv(
- "AZURE_OPENAI_AD_TOKEN"
- )
- values["openai_api_base"] = values["openai_api_base"] or os.getenv(
- "OPENAI_API_BASE"
- )
- values["openai_proxy"] = get_from_dict_or_env(
- values,
- "openai_proxy",
- "OPENAI_PROXY",
- default="",
- )
- values["openai_organization"] = (
- values["openai_organization"]
- or os.getenv("OPENAI_ORG_ID")
- or os.getenv("OPENAI_ORGANIZATION")
- )
- values["openai_api_version"] = values["openai_api_version"] or os.getenv(
- "OPENAI_API_VERSION"
- )
- values["openai_api_type"] = get_from_dict_or_env(
- values, "openai_api_type", "OPENAI_API_TYPE", default="azure"
- )
- try:
- import openai
- except ImportError:
- raise ImportError(
- "Could not import openai python package. "
- "Please install it with `pip install openai`."
- )
- if is_openai_v1():
- # For backwards compatibility. Before openai v1, no distinction was made
- # between azure_endpoint and base_url (openai_api_base).
- openai_api_base = values["openai_api_base"]
- if openai_api_base and values["validate_base_url"]:
- if "/openai" not in openai_api_base:
- values["openai_api_base"] = (
- values["openai_api_base"].rstrip("/") + "/openai"
- )
- warnings.warn(
- "As of openai>=1.0.0, Azure endpoints should be specified via "
- f"the `azure_endpoint` param not `openai_api_base` "
- f"(or alias `base_url`). Updating `openai_api_base` from "
- f"{openai_api_base} to {values['openai_api_base']}."
- )
- if values["deployment_name"]:
- warnings.warn(
- "As of openai>=1.0.0, if `deployment_name` (or alias "
- "`azure_deployment`) is specified then "
- "`openai_api_base` (or alias `base_url`) should not be. "
- "Instead use `deployment_name` (or alias `azure_deployment`) "
- "and `azure_endpoint`."
- )
- if values["deployment_name"] not in values["openai_api_base"]:
- warnings.warn(
- "As of openai>=1.0.0, if `openai_api_base` "
- "(or alias `base_url`) is specified it is expected to be "
- "of the form "
- "https://example-resource.azure.openai.com/openai/deployments/example-deployment. " # noqa: E501
- f"Updating {openai_api_base} to "
- f"{values['openai_api_base']}."
- )
- values["openai_api_base"] += (
- "/deployments/" + values["deployment_name"]
- )
- values["deployment_name"] = None
- client_params = {
- "api_version": values["openai_api_version"],
- "azure_endpoint": values["azure_endpoint"],
- "azure_deployment": values["deployment_name"],
- "api_key": values["openai_api_key"],
- "azure_ad_token": values["azure_ad_token"],
- "azure_ad_token_provider": values["azure_ad_token_provider"],
- "organization": values["openai_organization"],
- "base_url": values["openai_api_base"],
- "timeout": values["request_timeout"],
- "max_retries": values["max_retries"],
- "default_headers": {
- **(values["default_headers"] or {}),
- "User-Agent": "langchain-comm-python-azure-openai",
- },
- "default_query": values["default_query"],
- "http_client": values["http_client"],
- }
- values["client"] = openai.AzureOpenAI(**client_params).completions
-
- azure_ad_async_token_provider = values["azure_ad_async_token_provider"]
-
- if azure_ad_async_token_provider:
- client_params["azure_ad_token_provider"] = azure_ad_async_token_provider
-
- values["async_client"] = openai.AsyncAzureOpenAI(
- **client_params
- ).completions
-
- else:
- values["client"] = openai.Completion
-
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- return {
- **{"deployment_name": self.deployment_name},
- **super()._identifying_params,
- }
-
- @property
- def _invocation_params(self) -> Dict[str, Any]:
- if is_openai_v1():
- openai_params = {"model": self.deployment_name}
- else:
- openai_params = {
- "engine": self.deployment_name,
- "api_type": self.openai_api_type,
- "api_version": self.openai_api_version,
- }
- return {**openai_params, **super()._invocation_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "azure"
-
- @property
- def lc_attributes(self) -> Dict[str, Any]:
- return {
- "openai_api_type": self.openai_api_type,
- "openai_api_version": self.openai_api_version,
- }
-
-
-@deprecated(
- since="0.0.1",
- removal="1.0",
- alternative_import="langchain_openai.ChatOpenAI",
-)
-class OpenAIChat(BaseLLM):
- """OpenAI Chat large language models.
-
- To use, you should have the ``openai`` python package installed, and the
- environment variable ``OPENAI_API_KEY`` set with your API key.
-
- Any parameters that are valid to be passed to the openai.create call can be passed
- in, even if not explicitly saved on this class.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import OpenAIChat
- openaichat = OpenAIChat(model_name="gpt-3.5-turbo")
- """
-
- client: Any = Field(default=None, exclude=True) #: :meta private:
- async_client: Any = Field(default=None, exclude=True) #: :meta private:
- model_name: str = "gpt-3.5-turbo"
- """Model name to use."""
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `create` call not explicitly specified."""
- # When updating this to use a SecretStr
- # Check for classes that derive from this class (as some of them
- # may assume openai_api_key is a str)
- openai_api_key: Optional[str] = Field(default=None, alias="api_key")
- """Automatically inferred from env var `OPENAI_API_KEY` if not provided."""
- openai_api_base: Optional[str] = Field(default=None, alias="base_url")
- """Base URL path for API requests, leave blank if not using a proxy or service
- emulator."""
- # to support explicit proxy for OpenAI
- openai_proxy: Optional[str] = None
- max_retries: int = 6
- """Maximum number of retries to make when generating."""
- prefix_messages: List = Field(default_factory=list)
- """Series of messages for Chat input."""
- streaming: bool = False
- """Whether to stream the results or not."""
- allowed_special: Union[Literal["all"], AbstractSet[str]] = set()
- """Set of special tokens that are allowed。"""
- disallowed_special: Union[Literal["all"], Collection[str]] = "all"
- """Set of special tokens that are not allowed。"""
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = {field.alias for field in get_fields(cls).values()}
-
- extra = values.get("model_kwargs", {})
- for field_name in list(values):
- if field_name not in all_required_field_names:
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- extra[field_name] = values.pop(field_name)
- values["model_kwargs"] = extra
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- openai_api_key = get_from_dict_or_env(
- values, "openai_api_key", "OPENAI_API_KEY"
- )
- openai_api_base = get_from_dict_or_env(
- values,
- "openai_api_base",
- "OPENAI_API_BASE",
- default="",
- )
- openai_proxy = get_from_dict_or_env(
- values,
- "openai_proxy",
- "OPENAI_PROXY",
- default="",
- )
- openai_organization = get_from_dict_or_env(
- values, "openai_organization", "OPENAI_ORGANIZATION", default=""
- )
- try:
- import openai
-
- openai.api_key = openai_api_key
- if openai_api_base:
- openai.api_base = openai_api_base
- if openai_organization:
- openai.organization = openai_organization
- if openai_proxy:
- openai.proxy = {"http": openai_proxy, "https": openai_proxy}
- except ImportError:
- raise ImportError(
- "Could not import openai python package. "
- "Please install it with `pip install openai`."
- )
- try:
- values["client"] = openai.ChatCompletion
- except AttributeError:
- raise ValueError(
- "`openai` has no `ChatCompletion` attribute, this is likely "
- "due to an old version of the openai package. Try upgrading it "
- "with `pip install --upgrade openai`."
- )
- warnings.warn(
- "You are trying to use a chat model. This way of initializing it is "
- "no longer supported. Instead, please use: "
- "`from langchain_community.chat_models import ChatOpenAI`"
- )
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling OpenAI API."""
- return self.model_kwargs
-
- def _get_chat_params(
- self, prompts: List[str], stop: Optional[List[str]] = None
- ) -> Tuple:
- if len(prompts) > 1:
- raise ValueError(
- f"OpenAIChat currently only supports single prompt, got {prompts}"
- )
- messages = self.prefix_messages + [{"role": "user", "content": prompts[0]}]
- params: Dict[str, Any] = {**{"model": self.model_name}, **self._default_params}
- if stop is not None:
- if "stop" in params:
- raise ValueError("`stop` found in both the input and default params.")
- params["stop"] = stop
- if params.get("max_tokens") == -1:
- # for ChatGPT api, omitting max_tokens is equivalent to having no limit
- del params["max_tokens"]
- return messages, params
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- messages, params = self._get_chat_params([prompt], stop)
- params = {**params, **kwargs, "stream": True}
- for stream_resp in completion_with_retry(
- self, messages=messages, run_manager=run_manager, **params
- ):
- if not isinstance(stream_resp, dict):
- stream_resp = stream_resp.dict()
- token = stream_resp["choices"][0]["delta"].get("content", "")
- chunk = GenerationChunk(text=token)
- if run_manager:
- run_manager.on_llm_new_token(token, chunk=chunk)
- yield chunk
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- messages, params = self._get_chat_params([prompt], stop)
- params = {**params, **kwargs, "stream": True}
- async for stream_resp in await acompletion_with_retry(
- self, messages=messages, run_manager=run_manager, **params
- ):
- if not isinstance(stream_resp, dict):
- stream_resp = stream_resp.dict()
- token = stream_resp["choices"][0]["delta"].get("content", "")
- chunk = GenerationChunk(text=token)
- if run_manager:
- await run_manager.on_llm_new_token(token, chunk=chunk)
- yield chunk
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- if self.streaming:
- generation: Optional[GenerationChunk] = None
- for chunk in self._stream(prompts[0], stop, run_manager, **kwargs):
- if generation is None:
- generation = chunk
- else:
- generation += chunk
- assert generation is not None
- return LLMResult(generations=[[generation]])
-
- messages, params = self._get_chat_params(prompts, stop)
- params = {**params, **kwargs}
- full_response = completion_with_retry(
- self, messages=messages, run_manager=run_manager, **params
- )
- if not isinstance(full_response, dict):
- full_response = full_response.dict()
- llm_output = {
- "token_usage": full_response["usage"],
- "model_name": self.model_name,
- }
- return LLMResult(
- generations=[
- [Generation(text=full_response["choices"][0]["message"]["content"])]
- ],
- llm_output=llm_output,
- )
-
- async def _agenerate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- if self.streaming:
- generation: Optional[GenerationChunk] = None
- async for chunk in self._astream(prompts[0], stop, run_manager, **kwargs):
- if generation is None:
- generation = chunk
- else:
- generation += chunk
- assert generation is not None
- return LLMResult(generations=[[generation]])
-
- messages, params = self._get_chat_params(prompts, stop)
- params = {**params, **kwargs}
- full_response = await acompletion_with_retry(
- self, messages=messages, run_manager=run_manager, **params
- )
- if not isinstance(full_response, dict):
- full_response = full_response.dict()
- llm_output = {
- "token_usage": full_response["usage"],
- "model_name": self.model_name,
- }
- return LLMResult(
- generations=[
- [Generation(text=full_response["choices"][0]["message"]["content"])]
- ],
- llm_output=llm_output,
- )
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {**{"model_name": self.model_name}, **self._default_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "openai-chat"
-
- def get_token_ids(self, text: str) -> List[int]:
- """Get the token IDs using the tiktoken package."""
- # tiktoken NOT supported for Python < 3.8
- if sys.version_info[1] < 8:
- return super().get_token_ids(text)
- try:
- import tiktoken
- except ImportError:
- raise ImportError(
- "Could not import tiktoken python package. "
- "This is needed in order to calculate get_num_tokens. "
- "Please install it with `pip install tiktoken`."
- )
-
- enc = tiktoken.encoding_for_model(self.model_name)
- return enc.encode(
- text,
- allowed_special=self.allowed_special,
- disallowed_special=self.disallowed_special,
- )
diff --git a/libs/community/langchain_community/llms/openllm.py b/libs/community/langchain_community/llms/openllm.py
deleted file mode 100644
index 69c3944029..0000000000
--- a/libs/community/langchain_community/llms/openllm.py
+++ /dev/null
@@ -1,42 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Dict
-
-from langchain_community.llms.openai import BaseOpenAI
-from langchain_community.utils.openai import is_openai_v1
-
-
-class OpenLLM(BaseOpenAI):
- """OpenAI's compatible API client for OpenLLM server
-
- .. versionchanged:: 0.2.11
-
- Changed in 0.2.11 to support OpenLLM 0.6. Now behaves similar to OpenAI wrapper.
- """
-
- @classmethod
- def is_lc_serializable(cls) -> bool:
- return False
-
- @property
- def _invocation_params(self) -> Dict[str, Any]:
- """Get the parameters used to invoke the model."""
-
- params: Dict[str, Any] = {
- "model": self.model_name,
- **self._default_params,
- "logit_bias": None,
- }
- if not is_openai_v1():
- params.update(
- {
- "api_key": self.openai_api_key,
- "api_base": self.openai_api_base,
- }
- )
-
- return params
-
- @property
- def _llm_type(self) -> str:
- return "openllm"
diff --git a/libs/community/langchain_community/llms/openlm.py b/libs/community/langchain_community/llms/openlm.py
deleted file mode 100644
index 1601a3bd06..0000000000
--- a/libs/community/langchain_community/llms/openlm.py
+++ /dev/null
@@ -1,32 +0,0 @@
-from typing import Any, Dict
-
-from langchain_core.utils import pre_init
-
-from langchain_community.llms.openai import BaseOpenAI
-
-
-class OpenLM(BaseOpenAI):
- """OpenLM models."""
-
- @classmethod
- def is_lc_serializable(cls) -> bool:
- return False
-
- @property
- def _invocation_params(self) -> Dict[str, Any]:
- return {**{"model": self.model_name}, **super()._invocation_params}
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- try:
- import openlm
-
- values["client"] = openlm.Completion
- except ImportError:
- raise ImportError(
- "Could not import openlm python package. "
- "Please install it with `pip install openlm`."
- )
- if values["streaming"]:
- raise ValueError("Streaming not supported with openlm")
- return values
diff --git a/libs/community/langchain_community/llms/outlines.py b/libs/community/langchain_community/llms/outlines.py
deleted file mode 100644
index 25be8dcb4a..0000000000
--- a/libs/community/langchain_community/llms/outlines.py
+++ /dev/null
@@ -1,320 +0,0 @@
-from __future__ import annotations
-
-import importlib.util
-import logging
-import platform
-from typing import Any, Callable, Dict, Iterator, List, Literal, Optional, Tuple, Union
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from pydantic import BaseModel, Field, model_validator
-
-logger = logging.getLogger(__name__)
-
-
-class Outlines(LLM):
- """LLM wrapper for the Outlines library."""
-
- client: Any = None # :meta private:
-
- model: str
- """Identifier for the model to use with Outlines.
-
- The model identifier should be a string specifying:
- - A Hugging Face model name (e.g., "meta-llama/Llama-2-7b-chat-hf")
- - A local path to a model
- - For GGUF models, the format is "repo_id/file_name"
- (e.g., "TheBloke/Llama-2-7B-Chat-GGUF/llama-2-7b-chat.Q4_K_M.gguf")
-
- Examples:
- - "TheBloke/Llama-2-7B-Chat-GGUF/llama-2-7b-chat.Q4_K_M.gguf"
- - "meta-llama/Llama-2-7b-chat-hf"
- """
-
- backend: Literal[
- "llamacpp", "transformers", "transformers_vision", "vllm", "mlxlm"
- ] = "transformers"
- """Specifies the backend to use for the model.
-
- Supported backends are:
- - "llamacpp": For GGUF models using llama.cpp
- - "transformers": For Hugging Face Transformers models (default)
- - "transformers_vision": For vision-language models (e.g., LLaVA)
- - "vllm": For models using the vLLM library
- - "mlxlm": For models using the MLX framework
-
- Note: Ensure you have the necessary dependencies installed for the chosen backend.
- The system will attempt to import required packages and may raise an ImportError
- if they are not available.
- """
-
- max_tokens: int = 256
- """The maximum number of tokens to generate."""
-
- stop: Optional[List[str]] = None
- """A list of strings to stop generation when encountered."""
-
- streaming: bool = True
- """Whether to stream the results, token by token."""
-
- regex: Optional[str] = None
- r"""Regular expression for structured generation.
-
- If provided, Outlines will guarantee that the generated text matches this regex.
- This can be useful for generating structured outputs like IP addresses, dates, etc.
-
- Example: (valid IP address)
- regex = r"((25[0-5]|2[0-4]\d|[01]?\d\d?)\.){3}(25[0-5]|2[0-4]\d|[01]?\d\d?)"
-
- Note: Computing the regex index can take some time, so it's recommended to reuse
- the same regex for multiple generations if possible.
-
- For more details, see: https://dottxt-ai.github.io/outlines/reference/generation/regex/
- """
-
- type_constraints: Optional[Union[type, str]] = None
- """Type constraints for structured generation.
-
- Restricts the output to valid Python types. Supported types include:
- int, float, bool, datetime.date, datetime.time, datetime.datetime.
-
- Example:
- type_constraints = int
-
- For more details, see: https://dottxt-ai.github.io/outlines/reference/generation/format/
- """
-
- json_schema: Optional[Union[BaseModel, Dict, Callable]] = None
- """Pydantic model, JSON Schema, or callable (function signature)
- for structured JSON generation.
-
- Outlines can generate JSON output that follows a specified structure,
- which is useful for:
- 1. Parsing the answer (e.g., with Pydantic), storing it, or returning it to a user.
- 2. Calling a function with the result.
-
- You can provide:
- - A Pydantic model
- - A JSON Schema (as a Dict)
- - A callable (function signature)
-
- The generated JSON will adhere to the specified structure.
-
- For more details, see: https://dottxt-ai.github.io/outlines/reference/generation/json/
- """
-
- grammar: Optional[str] = None
- """Context-free grammar for structured generation.
-
- If provided, Outlines will generate text that adheres to the specified grammar.
- The grammar should be defined in EBNF format.
-
- This can be useful for generating structured outputs like mathematical expressions,
- programming languages, or custom domain-specific languages.
-
- Example:
- grammar = '''
- ?start: expression
- ?expression: term (("+" | "-") term)*
- ?term: factor (("*" | "/") factor)*
- ?factor: NUMBER | "-" factor | "(" expression ")"
- %import common.NUMBER
- '''
-
- Note: Grammar-based generation is currently experimental and may have performance
- limitations. It uses greedy generation to mitigate these issues.
-
- For more details and examples, see:
- https://dottxt-ai.github.io/outlines/reference/generation/cfg/
- """
-
- custom_generator: Optional[Any] = None
- """Set your own outlines generator object to override the default behavior."""
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Additional parameters to pass to the underlying model.
-
- Example:
- model_kwargs = {"temperature": 0.8, "seed": 42}
- """
-
- @model_validator(mode="after")
- def validate_environment(self) -> "Outlines":
- """Validate that outlines is installed and create a model instance."""
- num_constraints = sum(
- [
- bool(self.regex),
- bool(self.type_constraints),
- bool(self.json_schema),
- bool(self.grammar),
- ]
- )
- if num_constraints > 1:
- raise ValueError(
- "Either none or exactly one of regex, type_constraints, "
- "json_schema, or grammar can be provided."
- )
- return self.build_client()
-
- def build_client(self) -> "Outlines":
- try:
- import outlines.models as models
- except ImportError:
- raise ImportError(
- "Could not import the Outlines library. "
- "Please install it with `pip install outlines`."
- )
-
- def check_packages_installed(
- packages: List[Union[str, Tuple[str, str]]],
- ) -> None:
- missing_packages = [
- pkg if isinstance(pkg, str) else pkg[0]
- for pkg in packages
- if importlib.util.find_spec(pkg[1] if isinstance(pkg, tuple) else pkg)
- is None
- ]
- if missing_packages:
- raise ImportError( # todo this is displaying wrong
- f"Missing packages: {', '.join(missing_packages)}. "
- "You can install them with:\n\n"
- f" pip install {' '.join(missing_packages)}"
- )
-
- if self.backend == "llamacpp":
- if ".gguf" in self.model:
- creator, repo_name, file_name = self.model.split("/", 2)
- repo_id = f"{creator}/{repo_name}"
- else: # todo add auto-file-selection if no file is given
- raise ValueError("GGUF file_name must be provided for llama.cpp.")
- check_packages_installed([("llama-cpp-python", "llama_cpp")])
- self.client = models.llamacpp(repo_id, file_name, **self.model_kwargs)
- elif self.backend == "transformers":
- check_packages_installed(["transformers", "torch", "datasets"])
- self.client = models.transformers(self.model, **self.model_kwargs)
- elif self.backend == "transformers_vision":
- check_packages_installed(
- [
- "transformers",
- "datasets",
- "torchvision",
- "PIL",
- "flash_attn",
- ]
- )
- from transformers import LlavaNextForConditionalGeneration
-
- if not hasattr(models, "transformers_vision"):
- raise ValueError(
- "transformers_vision backend is not supported, "
- "please install the correct outlines version."
- )
- self.client = models.transformers_vision(
- self.model,
- model_class=LlavaNextForConditionalGeneration,
- **self.model_kwargs,
- )
- elif self.backend == "vllm":
- if platform.system() == "Darwin":
- raise ValueError("vLLM backend is not supported on macOS.")
- check_packages_installed(["vllm"])
- self.client = models.vllm(self.model, **self.model_kwargs)
- elif self.backend == "mlxlm":
- check_packages_installed(["mlx"])
- self.client = models.mlxlm(self.model, **self.model_kwargs)
- else:
- raise ValueError(f"Unsupported backend: {self.backend}")
-
- return self
-
- @property
- def _llm_type(self) -> str:
- return "outlines"
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- return {
- "max_tokens": self.max_tokens,
- "stop_at": self.stop,
- **self.model_kwargs,
- }
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- return {
- "model": self.model,
- "backend": self.backend,
- "regex": self.regex,
- "type_constraints": self.type_constraints,
- "json_schema": self.json_schema,
- "grammar": self.grammar,
- **self._default_params,
- }
-
- @property
- def _generator(self) -> Any:
- from outlines import generate
-
- if self.custom_generator:
- return self.custom_generator
- if self.regex:
- return generate.regex(self.client, regex_str=self.regex)
- if self.type_constraints:
- return generate.format(self.client, python_type=self.type_constraints)
- if self.json_schema:
- return generate.json(self.client, schema_object=self.json_schema)
- if self.grammar:
- return generate.cfg(self.client, cfg_str=self.grammar)
- return generate.text(self.client)
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- params = {**self._default_params, **kwargs}
- if stop:
- params["stop_at"] = stop
-
- response = ""
- if self.streaming:
- for chunk in self._stream(
- prompt=prompt,
- stop=params["stop_at"],
- run_manager=run_manager,
- **params,
- ):
- response += chunk.text
- else:
- response = self._generator(prompt, **params)
- return response
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- params = {**self._default_params, **kwargs}
- if stop:
- params["stop_at"] = stop
-
- for token in self._generator.stream(prompt, **params):
- if run_manager:
- run_manager.on_llm_new_token(token)
- yield GenerationChunk(text=token)
-
- @property
- def tokenizer(self) -> Any:
- """Access the tokenizer for the underlying model.
-
- .encode() to tokenize text.
- .decode() to convert tokens back to text.
- """
- if hasattr(self.client, "tokenizer"):
- return self.client.tokenizer
- raise ValueError("Tokenizer not found")
diff --git a/libs/community/langchain_community/llms/pai_eas_endpoint.py b/libs/community/langchain_community/llms/pai_eas_endpoint.py
deleted file mode 100644
index 5f447bda4c..0000000000
--- a/libs/community/langchain_community/llms/pai_eas_endpoint.py
+++ /dev/null
@@ -1,239 +0,0 @@
-import json
-import logging
-from typing import Any, Dict, Iterator, List, Mapping, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import get_from_dict_or_env, pre_init
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-class PaiEasEndpoint(LLM):
- """Langchain LLM class to help to access eass llm service.
-
- To use this endpoint, must have a deployed eas chat llm service on PAI AliCloud.
- One can set the environment variable ``eas_service_url`` and ``eas_service_token``.
- The environment variables can set with your eas service url and service token.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms.pai_eas_endpoint import PaiEasEndpoint
- eas_chat_endpoint = PaiEasChatEndpoint(
- eas_service_url="your_service_url",
- eas_service_token="your_service_token"
- )
- """
-
- """PAI-EAS Service URL"""
- eas_service_url: str
-
- """PAI-EAS Service TOKEN"""
- eas_service_token: str
-
- """PAI-EAS Service Infer Params"""
- max_new_tokens: Optional[int] = 512
- temperature: Optional[float] = 0.95
- top_p: Optional[float] = 0.1
- top_k: Optional[int] = 0
- stop_sequences: Optional[List[str]] = None
-
- """Enable stream chat mode."""
- streaming: bool = False
-
- """Key/value arguments to pass to the model. Reserved for future use"""
- model_kwargs: Optional[dict] = None
-
- version: Optional[str] = "2.0"
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- values["eas_service_url"] = get_from_dict_or_env(
- values, "eas_service_url", "EAS_SERVICE_URL"
- )
- values["eas_service_token"] = get_from_dict_or_env(
- values, "eas_service_token", "EAS_SERVICE_TOKEN"
- )
-
- return values
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "pai_eas_endpoint"
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling Cohere API."""
- return {
- "max_new_tokens": self.max_new_tokens,
- "temperature": self.temperature,
- "top_k": self.top_k,
- "top_p": self.top_p,
- "stop_sequences": [],
- }
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- "eas_service_url": self.eas_service_url,
- "eas_service_token": self.eas_service_token,
- **_model_kwargs,
- }
-
- def _invocation_params(
- self, stop_sequences: Optional[List[str]], **kwargs: Any
- ) -> dict:
- params = self._default_params
- if self.stop_sequences is not None and stop_sequences is not None:
- raise ValueError("`stop` found in both the input and default params.")
- elif self.stop_sequences is not None:
- params["stop"] = self.stop_sequences
- else:
- params["stop"] = stop_sequences
- if self.model_kwargs:
- params.update(self.model_kwargs)
- return {**params, **kwargs}
-
- @staticmethod
- def _process_response(
- response: Any, stop: Optional[List[str]], version: Optional[str]
- ) -> str:
- if version == "1.0":
- text = response
- else:
- text = response["response"]
-
- if stop:
- text = enforce_stop_tokens(text, stop)
- return "".join(text)
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- params = self._invocation_params(stop, **kwargs)
- prompt = prompt.strip()
- response = None
- try:
- if self.streaming:
- completion = ""
- for chunk in self._stream(prompt, stop, run_manager, **params):
- completion += chunk.text
- return completion
- else:
- response = self._call_eas(prompt, params)
- _stop = params.get("stop")
- return self._process_response(response, _stop, self.version)
- except Exception as error:
- raise ValueError(f"Error raised by the service: {error}")
-
- def _call_eas(self, prompt: str = "", params: Dict = {}) -> Any:
- """Generate text from the eas service."""
- headers = {
- "Content-Type": "application/json",
- "Authorization": f"{self.eas_service_token}",
- }
- if self.version == "1.0":
- body = {
- "input_ids": f"{prompt}",
- }
- else:
- body = {
- "prompt": f"{prompt}",
- }
-
- # add params to body
- for key, value in params.items():
- body[key] = value
-
- # make request
- response = requests.post(self.eas_service_url, headers=headers, json=body)
-
- if response.status_code != 200:
- raise Exception(
- f"Request failed with status code {response.status_code}"
- f" and message {response.text}"
- )
-
- try:
- return json.loads(response.text)
- except Exception as e:
- if isinstance(e, json.decoder.JSONDecodeError):
- return response.text
- raise e
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- invocation_params = self._invocation_params(stop, **kwargs)
-
- headers = {
- "User-Agent": "Test Client",
- "Authorization": f"{self.eas_service_token}",
- }
-
- if self.version == "1.0":
- pload = {"input_ids": prompt, **invocation_params}
- response = requests.post(
- self.eas_service_url, headers=headers, json=pload, stream=True
- )
-
- res = GenerationChunk(text=response.text)
-
- if run_manager:
- run_manager.on_llm_new_token(res.text)
-
- # yield text, if any
- yield res
- else:
- pload = {"prompt": prompt, "use_stream_chat": "True", **invocation_params}
-
- response = requests.post(
- self.eas_service_url, headers=headers, json=pload, stream=True
- )
-
- for chunk in response.iter_lines(
- chunk_size=8192, decode_unicode=False, delimiter=b"\0"
- ):
- if chunk:
- data = json.loads(chunk.decode("utf-8"))
- output = data["response"]
- # identify stop sequence in generated text, if any
- stop_seq_found: Optional[str] = None
- for stop_seq in invocation_params["stop"]:
- if stop_seq in output:
- stop_seq_found = stop_seq
-
- # identify text to yield
- text: Optional[str] = None
- if stop_seq_found:
- text = output[: output.index(stop_seq_found)]
- else:
- text = output
-
- # yield text, if any
- if text:
- res = GenerationChunk(text=text)
- if run_manager:
- run_manager.on_llm_new_token(res.text)
- yield res
-
- # break if stop sequence found
- if stop_seq_found:
- break
diff --git a/libs/community/langchain_community/llms/petals.py b/libs/community/langchain_community/llms/petals.py
deleted file mode 100644
index 7210037c6d..0000000000
--- a/libs/community/langchain_community/llms/petals.py
+++ /dev/null
@@ -1,154 +0,0 @@
-import logging
-from typing import Any, Dict, List, Mapping, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from langchain_core.utils.pydantic import get_fields
-from pydantic import ConfigDict, Field, SecretStr, model_validator
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-class Petals(LLM):
- """Petals Bloom models.
-
- To use, you should have the ``petals`` python package installed, and the
- environment variable ``HUGGINGFACE_API_KEY`` set with your API key.
-
- Any parameters that are valid to be passed to the call can be passed
- in, even if not explicitly saved on this class.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import petals
- petals = Petals()
-
- """
-
- client: Any = None
- """The client to use for the API calls."""
-
- tokenizer: Any = None
- """The tokenizer to use for the API calls."""
-
- model_name: str = "bigscience/bloom-petals"
- """The model to use."""
-
- temperature: float = 0.7
- """What sampling temperature to use"""
-
- max_new_tokens: int = 256
- """The maximum number of new tokens to generate in the completion."""
-
- top_p: float = 0.9
- """The cumulative probability for top-p sampling."""
-
- top_k: Optional[int] = None
- """The number of highest probability vocabulary tokens
- to keep for top-k-filtering."""
-
- do_sample: bool = True
- """Whether or not to use sampling; use greedy decoding otherwise."""
-
- max_length: Optional[int] = None
- """The maximum length of the sequence to be generated."""
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `create` call
- not explicitly specified."""
-
- huggingface_api_key: Optional[SecretStr] = None
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = {field.alias for field in get_fields(cls).values()}
-
- extra = values.get("model_kwargs", {})
- for field_name in list(values):
- if field_name not in all_required_field_names:
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- logger.warning(
- f"""WARNING! {field_name} is not default parameter.
- {field_name} was transferred to model_kwargs.
- Please confirm that {field_name} is what you intended."""
- )
- extra[field_name] = values.pop(field_name)
- values["model_kwargs"] = extra
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- huggingface_api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "huggingface_api_key", "HUGGINGFACE_API_KEY")
- )
- try:
- from petals import AutoDistributedModelForCausalLM
- from transformers import AutoTokenizer
-
- model_name = values["model_name"]
- values["tokenizer"] = AutoTokenizer.from_pretrained(model_name)
- values["client"] = AutoDistributedModelForCausalLM.from_pretrained(
- model_name
- )
- values["huggingface_api_key"] = huggingface_api_key.get_secret_value()
-
- except ImportError:
- raise ImportError(
- "Could not import transformers or petals python package."
- "Please install with `pip install -U transformers petals`."
- )
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling Petals API."""
- normal_params = {
- "temperature": self.temperature,
- "max_new_tokens": self.max_new_tokens,
- "top_p": self.top_p,
- "top_k": self.top_k,
- "do_sample": self.do_sample,
- "max_length": self.max_length,
- }
- return {**normal_params, **self.model_kwargs}
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {**{"model_name": self.model_name}, **self._default_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "petals"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call the Petals API."""
- params = self._default_params
- params = {**params, **kwargs}
- inputs = self.tokenizer(prompt, return_tensors="pt")["input_ids"]
- outputs = self.client.generate(inputs, **params)
- text = self.tokenizer.decode(outputs[0])
- if stop is not None:
- # I believe this is required since the stop tokens
- # are not enforced by the model parameters
- text = enforce_stop_tokens(text, stop)
- return text
diff --git a/libs/community/langchain_community/llms/pipelineai.py b/libs/community/langchain_community/llms/pipelineai.py
deleted file mode 100644
index 8d0f6e1579..0000000000
--- a/libs/community/langchain_community/llms/pipelineai.py
+++ /dev/null
@@ -1,121 +0,0 @@
-import logging
-from typing import Any, Dict, List, Mapping, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import (
- BaseModel,
- ConfigDict,
- Field,
- SecretStr,
- model_validator,
-)
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-class PipelineAI(LLM, BaseModel):
- """PipelineAI large language models.
-
- To use, you should have the ``pipeline-ai`` python package installed,
- and the environment variable ``PIPELINE_API_KEY`` set with your API key.
-
- Any parameters that are valid to be passed to the call can be passed
- in, even if not explicitly saved on this class.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import PipelineAI
- pipeline = PipelineAI(pipeline_key="")
- """
-
- pipeline_key: str = ""
- """The id or tag of the target pipeline"""
-
- pipeline_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any pipeline parameters valid for `create` call not
- explicitly specified."""
-
- pipeline_api_key: Optional[SecretStr] = None
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = set(list(cls.model_fields.keys()))
-
- extra = values.get("pipeline_kwargs", {})
- for field_name in list(values):
- if field_name not in all_required_field_names:
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- logger.warning(
- f"""{field_name} was transferred to pipeline_kwargs.
- Please confirm that {field_name} is what you intended."""
- )
- extra[field_name] = values.pop(field_name)
- values["pipeline_kwargs"] = extra
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- pipeline_api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "pipeline_api_key", "PIPELINE_API_KEY")
- )
- values["pipeline_api_key"] = pipeline_api_key
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- **{"pipeline_key": self.pipeline_key},
- **{"pipeline_kwargs": self.pipeline_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "pipeline_ai"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call to Pipeline Cloud endpoint."""
- try:
- from pipeline import PipelineCloud
- except ImportError:
- raise ImportError(
- "Could not import pipeline-ai python package. "
- "Please install it with `pip install pipeline-ai`."
- )
- client = PipelineCloud(token=self.pipeline_api_key.get_secret_value()) # type: ignore[union-attr]
- params = self.pipeline_kwargs or {}
- params = {**params, **kwargs}
-
- run = client.run_pipeline(self.pipeline_key, [prompt, params])
- try:
- text = run.result_preview[0][0]
- except AttributeError:
- raise AttributeError(
- f"A pipeline run should have a `result_preview` attribute."
- f"Run was: {run}"
- )
- if stop is not None:
- # I believe this is required since the stop tokens
- # are not enforced by the pipeline parameters
- text = enforce_stop_tokens(text, stop)
- return text
diff --git a/libs/community/langchain_community/llms/predibase.py b/libs/community/langchain_community/llms/predibase.py
deleted file mode 100644
index fbabdc04e7..0000000000
--- a/libs/community/langchain_community/llms/predibase.py
+++ /dev/null
@@ -1,218 +0,0 @@
-import os
-from typing import Any, Dict, List, Mapping, Optional, Union
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from pydantic import Field, SecretStr
-
-
-class Predibase(LLM):
- """Use your Predibase models with Langchain.
-
- To use, you should have the ``predibase`` python package installed,
- and have your Predibase API key.
-
- The `model` parameter is the Predibase "serverless" base_model ID
- (see https://docs.predibase.com/user-guide/inference/models for the catalog).
-
- An optional `adapter_id` parameter is the Predibase ID or HuggingFace ID of a
- fine-tuned LLM adapter, whose base model is the `model` parameter; the
- fine-tuned adapter must be compatible with its base model;
- otherwise, an error is raised. If the fine-tuned adapter is hosted at Predibase,
- then `adapter_version` in the adapter repository must be specified.
-
- An optional `predibase_sdk_version` parameter defaults to latest SDK version.
- """
-
- model: str
- predibase_api_key: SecretStr
- predibase_sdk_version: Optional[str] = None
- adapter_id: Optional[str] = None
- adapter_version: Optional[int] = None
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- default_options_for_generation: dict = Field(
- {
- "max_new_tokens": 256,
- "temperature": 0.1,
- }
- )
-
- @property
- def _llm_type(self) -> str:
- return "predibase"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- options: Dict[str, Union[str, float]] = {
- **self.default_options_for_generation,
- **(self.model_kwargs or {}),
- **(kwargs or {}),
- }
- if self._is_deprecated_sdk_version():
- try:
- from predibase import PredibaseClient
- from predibase.pql import get_session
- from predibase.pql.api import (
- ServerResponseError,
- Session,
- )
- from predibase.resource.llm.interface import (
- HuggingFaceLLM,
- LLMDeployment,
- )
- from predibase.resource.llm.response import GeneratedResponse
- from predibase.resource.model import Model
-
- session: Session = get_session(
- token=self.predibase_api_key.get_secret_value(),
- gateway="https://api.app.predibase.com/v1",
- serving_endpoint="serving.app.predibase.com",
- )
- pc: PredibaseClient = PredibaseClient(session=session)
- except ImportError as e:
- raise ImportError(
- "Could not import Predibase Python package. "
- "Please install it with `pip install predibase`."
- ) from e
- except ValueError as e:
- raise ValueError("Your API key is not correct. Please try again") from e
-
- base_llm_deployment: LLMDeployment = pc.LLM(
- uri=f"pb://deployments/{self.model}"
- )
- result: GeneratedResponse
- if self.adapter_id:
- """
- Attempt to retrieve the fine-tuned adapter from a Predibase
- repository. If absent, then load the fine-tuned adapter
- from a HuggingFace repository.
- """
- adapter_model: Union[Model, HuggingFaceLLM]
- try:
- adapter_model = pc.get_model(
- name=self.adapter_id,
- version=self.adapter_version,
- model_id=None,
- )
- except ServerResponseError:
- # Predibase does not recognize the adapter ID (query HuggingFace).
- adapter_model = pc.LLM(uri=f"hf://{self.adapter_id}")
- result = base_llm_deployment.with_adapter(model=adapter_model).generate(
- prompt=prompt,
- options=options,
- )
- else:
- result = base_llm_deployment.generate(
- prompt=prompt,
- options=options,
- )
- return result.response
-
- from predibase import Predibase
-
- os.environ["PREDIBASE_GATEWAY"] = "https://api.app.predibase.com"
- predibase: Predibase = Predibase(
- api_token=self.predibase_api_key.get_secret_value()
- )
-
- import requests
- from lorax.client import Client as LoraxClient
- from lorax.errors import GenerationError
- from lorax.types import Response
-
- lorax_client: LoraxClient = predibase.deployments.client(
- deployment_ref=self.model
- )
-
- response: Response
- if self.adapter_id:
- """
- Attempt to retrieve the fine-tuned adapter from a Predibase repository.
- If absent, then load the fine-tuned adapter from a HuggingFace repository.
- """
- if self.adapter_version:
- # Since the adapter version is provided, query the Predibase repository.
- pb_adapter_id: str = f"{self.adapter_id}/{self.adapter_version}"
- options.pop(
- "api_token", None
- ) # The "api_token" is not used for Predibase-hosted models.
- try:
- response = lorax_client.generate(
- prompt=prompt,
- adapter_id=pb_adapter_id,
- **options,
- )
- except GenerationError as ge:
- raise ValueError(
- f"""An adapter with the ID "{pb_adapter_id}" cannot be \
-found in the Predibase repository of fine-tuned adapters."""
- ) from ge
- else:
- # The adapter version is omitted,
- # hence look for the adapter ID in the HuggingFace repository.
- try:
- response = lorax_client.generate(
- prompt=prompt,
- adapter_id=self.adapter_id,
- adapter_source="hub",
- **options,
- )
- except GenerationError as ge:
- raise ValueError(
- f"""Either an adapter with the ID "{self.adapter_id}" \
-cannot be found in a HuggingFace repository, or it is incompatible with the \
-base model (please make sure that the adapter configuration is consistent).
-"""
- ) from ge
- else:
- try:
- response = lorax_client.generate(
- prompt=prompt,
- **options,
- )
- except requests.JSONDecodeError as jde:
- raise ValueError(
- f"""An LLM with the deployment ID "{self.model}" cannot be found \
-at Predibase (please refer to \
-"https://docs.predibase.com/user-guide/inference/models" for the list of \
-supported models).
-"""
- ) from jde
- response_text = response.generated_text
-
- return response_text
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- **{"model_kwargs": self.model_kwargs},
- }
-
- def _is_deprecated_sdk_version(self) -> bool:
- try:
- import semantic_version
- from predibase.version import __version__ as current_version
- from semantic_version.base import Version
-
- sdk_semver_deprecated: Version = semantic_version.Version(
- version_string="2024.4.8"
- )
- actual_current_version: str = self.predibase_sdk_version or current_version
- sdk_semver_current: Version = semantic_version.Version(
- version_string=actual_current_version
- )
- return not (
- (sdk_semver_current > sdk_semver_deprecated)
- or ("+dev" in actual_current_version)
- )
- except ImportError as e:
- raise ImportError(
- "Could not import Predibase Python package. "
- "Please install it with `pip install semantic_version predibase`."
- ) from e
diff --git a/libs/community/langchain_community/llms/predictionguard.py b/libs/community/langchain_community/llms/predictionguard.py
deleted file mode 100644
index 01edbfa16d..0000000000
--- a/libs/community/langchain_community/llms/predictionguard.py
+++ /dev/null
@@ -1,166 +0,0 @@
-import logging
-from typing import Any, Dict, List, Optional, Union
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import BaseModel, ConfigDict, model_validator
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-@deprecated(
- since="0.3.28",
- removal="1.0",
- alternative_import="langchain_predictionguard.PredictionGuard",
-)
-class PredictionGuard(LLM):
- """Prediction Guard large language models.
-
- To use, you should have the ``predictionguard`` python package installed, and the
- environment variable ``PREDICTIONGUARD_API_KEY`` set with your API key, or pass
- it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- llm = PredictionGuard(
- model="Hermes-3-Llama-3.1-8B",
- predictionguard_api_key="your Prediction Guard API key",
- )
- """
-
- client: Any = None #: :meta private:
-
- model: Optional[str] = "Hermes-3-Llama-3.1-8B"
- """Model name to use."""
-
- max_tokens: Optional[int] = 256
- """Denotes the number of tokens to predict per generation."""
-
- temperature: Optional[float] = 0.75
- """A non-negative float that tunes the degree of randomness in generation."""
-
- top_p: Optional[float] = 0.1
- """A non-negative float that controls the diversity of the generated tokens."""
-
- top_k: Optional[int] = None
- """The diversity of the generated text based on top-k sampling."""
-
- stop: Optional[List[str]] = None
-
- predictionguard_input: Optional[Dict[str, Union[str, bool]]] = None
- """The input check to run over the prompt before sending to the LLM."""
-
- predictionguard_output: Optional[Dict[str, bool]] = None
- """The output check to run the LLM output against."""
-
- predictionguard_api_key: Optional[str] = None
- """Prediction Guard API key."""
-
- model_config = ConfigDict(extra="forbid")
-
- @model_validator(mode="before")
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that the api_key and python package exists in environment."""
- pg_api_key = get_from_dict_or_env(
- values, "predictionguard_api_key", "PREDICTIONGUARD_API_KEY"
- )
-
- try:
- from predictionguard import PredictionGuard
-
- values["client"] = PredictionGuard(
- api_key=pg_api_key,
- )
-
- except ImportError:
- raise ImportError(
- "Could not import predictionguard python package. "
- "Please install it with `pip install predictionguard`."
- )
-
- return values
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {"model": self.model}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "predictionguard"
-
- def _get_parameters(self, **kwargs: Any) -> Dict[str, Any]:
- # input kwarg conflicts with LanguageModelInput on BaseChatModel
- input = kwargs.pop("predictionguard_input", self.predictionguard_input)
- output = kwargs.pop("predictionguard_output", self.predictionguard_output)
-
- params = {
- **{
- "max_tokens": self.max_tokens,
- "temperature": self.temperature,
- "top_p": self.top_p,
- "top_k": self.top_k,
- "input": (
- input.model_dump() if isinstance(input, BaseModel) else input
- ),
- "output": (
- output.model_dump() if isinstance(output, BaseModel) else output
- ),
- },
- **kwargs,
- }
-
- return params
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Prediction Guard's model API.
- Args:
- prompt: The prompt to pass into the model.
- Returns:
- The string generated by the model.
- Example:
- .. code-block:: python
- response = llm.invoke("Tell me a joke.")
- """
-
- params = self._get_parameters(**kwargs)
-
- stops = None
- if self.stop is not None and stop is not None:
- raise ValueError("`stop` found in both the input and default params.")
- elif self.stop is not None:
- stops = self.stop
- else:
- stops = stop
-
- response = self.client.completions.create(
- model=self.model,
- prompt=prompt,
- **params,
- )
-
- for res in response["choices"]:
- if res.get("status", "").startswith("error: "):
- err_msg = res["status"].removeprefix("error: ")
- raise ValueError(f"Error from PredictionGuard API: {err_msg}")
-
- text = response["choices"][0]["text"]
-
- # If stop tokens are provided, Prediction Guard's endpoint returns them.
- # In order to make this consistent with other endpoints, we strip them.
- if stops:
- text = enforce_stop_tokens(text, stops)
-
- return text
diff --git a/libs/community/langchain_community/llms/promptlayer_openai.py b/libs/community/langchain_community/llms/promptlayer_openai.py
deleted file mode 100644
index 15456a7399..0000000000
--- a/libs/community/langchain_community/llms/promptlayer_openai.py
+++ /dev/null
@@ -1,232 +0,0 @@
-import datetime
-from typing import Any, List, Optional
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.outputs import LLMResult
-
-from langchain_community.llms.openai import OpenAI, OpenAIChat
-
-
-class PromptLayerOpenAI(OpenAI):
- """PromptLayer OpenAI large language models.
-
- To use, you should have the ``openai`` and ``promptlayer`` python
- package installed, and the environment variable ``OPENAI_API_KEY``
- and ``PROMPTLAYER_API_KEY`` set with your openAI API key and
- promptlayer key respectively.
-
- All parameters that can be passed to the OpenAI LLM can also
- be passed here. The PromptLayerOpenAI LLM adds two optional
-
- parameters:
- ``pl_tags``: List of strings to tag the request with.
- ``return_pl_id``: If True, the PromptLayer request ID will be
- returned in the ``generation_info`` field of the
- ``Generation`` object.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import PromptLayerOpenAI
- openai = PromptLayerOpenAI(model_name="gpt-3.5-turbo-instruct")
- """
-
- pl_tags: Optional[List[str]]
- return_pl_id: Optional[bool] = False
-
- @classmethod
- def is_lc_serializable(cls) -> bool:
- return False
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call OpenAI generate and then call PromptLayer API to log the request."""
- from promptlayer.utils import get_api_key, promptlayer_api_request
-
- request_start_time = datetime.datetime.now().timestamp()
- generated_responses = super()._generate(prompts, stop, run_manager)
- request_end_time = datetime.datetime.now().timestamp()
- for i in range(len(prompts)):
- prompt = prompts[i]
- generation = generated_responses.generations[i][0]
- resp = {
- "text": generation.text,
- "llm_output": generated_responses.llm_output,
- }
- params = {**self._identifying_params, **kwargs}
- pl_request_id = promptlayer_api_request(
- "langchain.PromptLayerOpenAI",
- "langchain",
- [prompt],
- params,
- self.pl_tags,
- resp,
- request_start_time,
- request_end_time,
- get_api_key(),
- return_pl_id=self.return_pl_id,
- )
- if self.return_pl_id:
- if generation.generation_info is None or not isinstance(
- generation.generation_info, dict
- ):
- generation.generation_info = {}
- generation.generation_info["pl_request_id"] = pl_request_id
- return generated_responses
-
- async def _agenerate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- from promptlayer.utils import get_api_key, promptlayer_api_request_async
-
- request_start_time = datetime.datetime.now().timestamp()
- generated_responses = await super()._agenerate(prompts, stop, run_manager)
- request_end_time = datetime.datetime.now().timestamp()
- for i in range(len(prompts)):
- prompt = prompts[i]
- generation = generated_responses.generations[i][0]
- resp = {
- "text": generation.text,
- "llm_output": generated_responses.llm_output,
- }
- params = {**self._identifying_params, **kwargs}
- pl_request_id = await promptlayer_api_request_async(
- "langchain.PromptLayerOpenAI.async",
- "langchain",
- [prompt],
- params,
- self.pl_tags,
- resp,
- request_start_time,
- request_end_time,
- get_api_key(),
- return_pl_id=self.return_pl_id,
- )
- if self.return_pl_id:
- if generation.generation_info is None or not isinstance(
- generation.generation_info, dict
- ):
- generation.generation_info = {}
- generation.generation_info["pl_request_id"] = pl_request_id
- return generated_responses
-
-
-class PromptLayerOpenAIChat(OpenAIChat):
- """PromptLayer OpenAI large language models.
-
- To use, you should have the ``openai`` and ``promptlayer`` python
- package installed, and the environment variable ``OPENAI_API_KEY``
- and ``PROMPTLAYER_API_KEY`` set with your openAI API key and
- promptlayer key respectively.
-
- All parameters that can be passed to the OpenAIChat LLM can also
- be passed here. The PromptLayerOpenAIChat adds two optional
-
- parameters:
- ``pl_tags``: List of strings to tag the request with.
- ``return_pl_id``: If True, the PromptLayer request ID will be
- returned in the ``generation_info`` field of the
- ``Generation`` object.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import PromptLayerOpenAIChat
- openaichat = PromptLayerOpenAIChat(model_name="gpt-3.5-turbo")
- """
-
- pl_tags: Optional[List[str]]
- return_pl_id: Optional[bool] = False
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call OpenAI generate and then call PromptLayer API to log the request."""
- from promptlayer.utils import get_api_key, promptlayer_api_request
-
- request_start_time = datetime.datetime.now().timestamp()
- generated_responses = super()._generate(prompts, stop, run_manager)
- request_end_time = datetime.datetime.now().timestamp()
- for i in range(len(prompts)):
- prompt = prompts[i]
- generation = generated_responses.generations[i][0]
- resp = {
- "text": generation.text,
- "llm_output": generated_responses.llm_output,
- }
- params = {**self._identifying_params, **kwargs}
- pl_request_id = promptlayer_api_request(
- "langchain.PromptLayerOpenAIChat",
- "langchain",
- [prompt],
- params,
- self.pl_tags,
- resp,
- request_start_time,
- request_end_time,
- get_api_key(),
- return_pl_id=self.return_pl_id,
- )
- if self.return_pl_id:
- if generation.generation_info is None or not isinstance(
- generation.generation_info, dict
- ):
- generation.generation_info = {}
- generation.generation_info["pl_request_id"] = pl_request_id
- return generated_responses
-
- async def _agenerate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- from promptlayer.utils import get_api_key, promptlayer_api_request_async
-
- request_start_time = datetime.datetime.now().timestamp()
- generated_responses = await super()._agenerate(prompts, stop, run_manager)
- request_end_time = datetime.datetime.now().timestamp()
- for i in range(len(prompts)):
- prompt = prompts[i]
- generation = generated_responses.generations[i][0]
- resp = {
- "text": generation.text,
- "llm_output": generated_responses.llm_output,
- }
- params = {**self._identifying_params, **kwargs}
- pl_request_id = await promptlayer_api_request_async(
- "langchain.PromptLayerOpenAIChat.async",
- "langchain",
- [prompt],
- params,
- self.pl_tags,
- resp,
- request_start_time,
- request_end_time,
- get_api_key(),
- return_pl_id=self.return_pl_id,
- )
- if self.return_pl_id:
- if generation.generation_info is None or not isinstance(
- generation.generation_info, dict
- ):
- generation.generation_info = {}
- generation.generation_info["pl_request_id"] = pl_request_id
- return generated_responses
diff --git a/libs/community/langchain_community/llms/replicate.py b/libs/community/langchain_community/llms/replicate.py
deleted file mode 100644
index f6c4e15ba6..0000000000
--- a/libs/community/langchain_community/llms/replicate.py
+++ /dev/null
@@ -1,232 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import TYPE_CHECKING, Any, Dict, Iterator, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from langchain_core.utils.pydantic import get_fields
-from pydantic import ConfigDict, Field, model_validator
-
-if TYPE_CHECKING:
- from replicate.prediction import Prediction
-
-logger = logging.getLogger(__name__)
-
-
-class Replicate(LLM):
- """Replicate models.
-
- To use, you should have the ``replicate`` python package installed,
- and the environment variable ``REPLICATE_API_TOKEN`` set with your API token.
- You can find your token here: https://replicate.com/account
-
- The model param is required, but any other model parameters can also
- be passed in with the format model_kwargs={model_param: value, ...}
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import Replicate
-
- replicate = Replicate(
- model=(
- "stability-ai/stable-diffusion: "
- "27b93a2413e7f36cd83da926f3656280b2931564ff050bf9575f1fdf9bcd7478",
- ),
- model_kwargs={"image_dimensions": "512x512"}
- )
- """
-
- model: str
- model_kwargs: Dict[str, Any] = Field(default_factory=dict, alias="input")
- replicate_api_token: Optional[str] = None
- prompt_key: Optional[str] = None
- version_obj: Any = Field(default=None, exclude=True)
- """Optionally pass in the model version object during initialization to avoid
- having to make an extra API call to retrieve it during streaming. NOTE: not
- serializable, is excluded from serialization.
- """
-
- streaming: bool = False
- """Whether to stream the results."""
-
- stop: List[str] = Field(default_factory=list)
- """Stop sequences to early-terminate generation."""
-
- model_config = ConfigDict(
- populate_by_name=True,
- extra="forbid",
- )
-
- @property
- def lc_secrets(self) -> Dict[str, str]:
- return {"replicate_api_token": "REPLICATE_API_TOKEN"}
-
- @classmethod
- def is_lc_serializable(cls) -> bool:
- return True
-
- @classmethod
- def get_lc_namespace(cls) -> List[str]:
- """Get the namespace of the langchain object."""
- return ["langchain", "llms", "replicate"]
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = {field for field in get_fields(cls).keys()}
-
- input = values.pop("input", {})
- if input:
- logger.warning(
- "Init param `input` is deprecated, please use `model_kwargs` instead."
- )
- extra = {**values.pop("model_kwargs", {}), **input}
- for field_name in list(values):
- if field_name not in all_required_field_names:
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- logger.warning(
- f"""{field_name} was transferred to model_kwargs.
- Please confirm that {field_name} is what you intended."""
- )
- extra[field_name] = values.pop(field_name)
- values["model_kwargs"] = extra
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- replicate_api_token = get_from_dict_or_env(
- values, "replicate_api_token", "REPLICATE_API_TOKEN"
- )
- values["replicate_api_token"] = replicate_api_token
- return values
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {
- "model": self.model,
- "model_kwargs": self.model_kwargs,
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of model."""
- return "replicate"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call to replicate endpoint."""
- if self.streaming:
- completion: Optional[str] = None
- for chunk in self._stream(
- prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- if completion is None:
- completion = chunk.text
- else:
- completion += chunk.text
- else:
- prediction = self._create_prediction(prompt, **kwargs)
- prediction.wait()
- if prediction.status == "failed":
- raise RuntimeError(prediction.error)
- if isinstance(prediction.output, str):
- completion = prediction.output
- else:
- completion = "".join(prediction.output)
- assert completion is not None
- stop_conditions = stop or self.stop
- for s in stop_conditions:
- if s in completion:
- completion = completion[: completion.find(s)]
- return completion
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- prediction = self._create_prediction(prompt, **kwargs)
- stop_conditions = stop or self.stop
- stop_condition_reached = False
- current_completion: str = ""
- for output in prediction.output_iterator():
- current_completion += output
- # test for stop conditions, if specified
- for s in stop_conditions:
- if s in current_completion:
- prediction.cancel()
- stop_condition_reached = True
- # Potentially some tokens that should still be yielded before ending
- # stream.
- stop_index = max(output.find(s), 0)
- output = output[:stop_index]
- if not output:
- break
- if output:
- if run_manager:
- run_manager.on_llm_new_token(
- output,
- verbose=self.verbose,
- )
- yield GenerationChunk(text=output)
- if stop_condition_reached:
- break
-
- def _create_prediction(self, prompt: str, **kwargs: Any) -> Prediction:
- try:
- import replicate as replicate_python
- except ImportError:
- raise ImportError(
- "Could not import replicate python package. "
- "Please install it with `pip install replicate`."
- )
-
- # get the model and version
- if self.version_obj is None:
- if ":" in self.model:
- model_str, version_str = self.model.split(":")
- model = replicate_python.models.get(model_str)
- self.version_obj = model.versions.get(version_str)
- else:
- model = replicate_python.models.get(self.model)
- self.version_obj = model.latest_version
-
- if self.prompt_key is None:
- # sort through the openapi schema to get the name of the first input
- input_properties = sorted(
- self.version_obj.openapi_schema["components"]["schemas"]["Input"][
- "properties"
- ].items(),
- key=lambda item: item[1].get("x-order", 0),
- )
-
- self.prompt_key = input_properties[0][0]
-
- input_: Dict = {
- self.prompt_key: prompt,
- **self.model_kwargs,
- **kwargs,
- }
-
- # if it's an official model
- if ":" not in self.model:
- return replicate_python.models.predictions.create(self.model, input=input_)
- else:
- return replicate_python.predictions.create(
- version=self.version_obj, input=input_
- )
diff --git a/libs/community/langchain_community/llms/rwkv.py b/libs/community/langchain_community/llms/rwkv.py
deleted file mode 100644
index 9273467c7f..0000000000
--- a/libs/community/langchain_community/llms/rwkv.py
+++ /dev/null
@@ -1,235 +0,0 @@
-"""RWKV models.
-
-Based on https://github.com/saharNooby/rwkv.cpp/blob/master/rwkv/chat_with_bot.py
- https://github.com/BlinkDL/ChatRWKV/blob/main/v2/chat.py
-"""
-
-from typing import Any, Dict, List, Mapping, Optional, Set
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import pre_init
-from pydantic import BaseModel, ConfigDict
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-
-class RWKV(LLM, BaseModel):
- """RWKV language models.
-
- To use, you should have the ``rwkv`` python package installed, the
- pre-trained model file, and the model's config information.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import RWKV
- model = RWKV(model="./models/rwkv-3b-fp16.bin", strategy="cpu fp32")
-
- # Simplest invocation
- response = model.invoke("Once upon a time, ")
- """
-
- model: str
- """Path to the pre-trained RWKV model file."""
-
- tokens_path: str
- """Path to the RWKV tokens file."""
-
- strategy: str = "cpu fp32"
- """Token context window."""
-
- rwkv_verbose: bool = True
- """Print debug information."""
-
- temperature: float = 1.0
- """The temperature to use for sampling."""
-
- top_p: float = 0.5
- """The top-p value to use for sampling."""
-
- penalty_alpha_frequency: float = 0.4
- """Positive values penalize new tokens based on their existing frequency
- in the text so far, decreasing the model's likelihood to repeat the same
- line verbatim.."""
-
- penalty_alpha_presence: float = 0.4
- """Positive values penalize new tokens based on whether they appear
- in the text so far, increasing the model's likelihood to talk about
- new topics.."""
-
- CHUNK_LEN: int = 256
- """Batch size for prompt processing."""
-
- max_tokens_per_generation: int = 256
- """Maximum number of tokens to generate."""
-
- client: Any = None #: :meta private:
-
- tokenizer: Any = None #: :meta private:
-
- pipeline: Any = None #: :meta private:
-
- model_tokens: Any = None #: :meta private:
-
- model_state: Any = None #: :meta private:
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {
- "verbose": self.verbose,
- "top_p": self.top_p,
- "temperature": self.temperature,
- "penalty_alpha_frequency": self.penalty_alpha_frequency,
- "penalty_alpha_presence": self.penalty_alpha_presence,
- "CHUNK_LEN": self.CHUNK_LEN,
- "max_tokens_per_generation": self.max_tokens_per_generation,
- }
-
- @staticmethod
- def _rwkv_param_names() -> Set[str]:
- """Get the identifying parameters."""
- return {
- "verbose",
- }
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that the python package exists in the environment."""
- try:
- import tokenizers
- except ImportError:
- raise ImportError(
- "Could not import tokenizers python package. "
- "Please install it with `pip install tokenizers`."
- )
- try:
- from rwkv.model import RWKV as RWKVMODEL
- from rwkv.utils import PIPELINE
-
- values["tokenizer"] = tokenizers.Tokenizer.from_file(values["tokens_path"])
-
- rwkv_keys = cls._rwkv_param_names()
- model_kwargs = {k: v for k, v in values.items() if k in rwkv_keys}
- model_kwargs["verbose"] = values["rwkv_verbose"]
- values["client"] = RWKVMODEL(
- values["model"], strategy=values["strategy"], **model_kwargs
- )
- values["pipeline"] = PIPELINE(values["client"], values["tokens_path"])
-
- except ImportError:
- raise ImportError(
- "Could not import rwkv python package. "
- "Please install it with `pip install rwkv`."
- )
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- "model": self.model,
- **self._default_params,
- **{k: v for k, v in self.__dict__.items() if k in RWKV._rwkv_param_names()},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return the type of llm."""
- return "rwkv"
-
- def run_rnn(self, _tokens: List[str], newline_adj: int = 0) -> Any:
- AVOID_REPEAT_TOKENS = []
- AVOID_REPEAT = ",:?!"
- for i in AVOID_REPEAT:
- dd = self.pipeline.encode(i)
- assert len(dd) == 1
- AVOID_REPEAT_TOKENS += dd
-
- tokens = [int(x) for x in _tokens]
- self.model_tokens += tokens
-
- out: Any = None
-
- while len(tokens) > 0:
- out, self.model_state = self.client.forward(
- tokens[: self.CHUNK_LEN], self.model_state
- )
- tokens = tokens[self.CHUNK_LEN :]
- END_OF_LINE = 187
- out[END_OF_LINE] += newline_adj # adjust \n probability
-
- if self.model_tokens[-1] in AVOID_REPEAT_TOKENS:
- out[self.model_tokens[-1]] = -999999999
- return out
-
- def rwkv_generate(self, prompt: str) -> str:
- self.model_state = None
- self.model_tokens = []
- logits = self.run_rnn(self.tokenizer.encode(prompt).ids)
- begin = len(self.model_tokens)
- out_last = begin
-
- occurrence: Dict = {}
-
- decoded = ""
- for i in range(self.max_tokens_per_generation):
- for n in occurrence:
- logits[n] -= (
- self.penalty_alpha_presence
- + occurrence[n] * self.penalty_alpha_frequency
- )
- token = self.pipeline.sample_logits(
- logits, temperature=self.temperature, top_p=self.top_p
- )
-
- END_OF_TEXT = 0
- if token == END_OF_TEXT:
- break
- if token not in occurrence:
- occurrence[token] = 1
- else:
- occurrence[token] += 1
-
- logits = self.run_rnn([token])
- xxx = self.tokenizer.decode(self.model_tokens[out_last:])
- if "\ufffd" not in xxx: # avoid utf-8 display issues
- decoded += xxx
- out_last = begin + i + 1
- if i >= self.max_tokens_per_generation - 100:
- break
-
- return decoded
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- r"""RWKV generation
-
- Args:
- prompt: The prompt to pass into the model.
- stop: A list of strings to stop generation when encountered.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- prompt = "Once upon a time, "
- response = model.invoke(prompt, n_predict=55)
- """
- text = self.rwkv_generate(prompt)
-
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
- return text
diff --git a/libs/community/langchain_community/llms/sagemaker_endpoint.py b/libs/community/langchain_community/llms/sagemaker_endpoint.py
deleted file mode 100644
index e78d265441..0000000000
--- a/libs/community/langchain_community/llms/sagemaker_endpoint.py
+++ /dev/null
@@ -1,377 +0,0 @@
-"""Sagemaker InvokeEndpoint API."""
-
-import io
-import json
-from abc import abstractmethod
-from typing import Any, Dict, Generic, Iterator, List, Mapping, Optional, TypeVar, Union
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import pre_init
-from pydantic import ConfigDict
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-INPUT_TYPE = TypeVar("INPUT_TYPE", bound=Union[str, List[str]])
-OUTPUT_TYPE = TypeVar("OUTPUT_TYPE", bound=Union[str, List[List[float]], Iterator])
-
-
-class LineIterator:
- """Parse the byte stream input.
-
- The output of the model will be in the following format:
-
- b'{"outputs": [" a"]}\n'
- b'{"outputs": [" challenging"]}\n'
- b'{"outputs": [" problem"]}\n'
- ...
-
- While usually each PayloadPart event from the event stream will
- contain a byte array with a full json, this is not guaranteed
- and some of the json objects may be split acrossPayloadPart events.
-
- For example:
-
- {'PayloadPart': {'Bytes': b'{"outputs": '}}
- {'PayloadPart': {'Bytes': b'[" problem"]}\n'}}
-
-
- This class accounts for this by concatenating bytes written via the 'write' function
- and then exposing a method which will return lines (ending with a '\n' character)
- within the buffer via the 'scan_lines' function.
- It maintains the position of the last read position to ensure
- that previous bytes are not exposed again.
-
- For more details see:
- https://aws.amazon.com/blogs/machine-learning/elevating-the-generative-ai-experience-introducing-streaming-support-in-amazon-sagemaker-hosting/
- """
-
- def __init__(self, stream: Any) -> None:
- self.byte_iterator = iter(stream)
- self.buffer = io.BytesIO()
- self.read_pos = 0
-
- def __iter__(self) -> "LineIterator":
- return self
-
- def __next__(self) -> Any:
- while True:
- self.buffer.seek(self.read_pos)
- line = self.buffer.readline()
- if line and line[-1] == ord("\n"):
- self.read_pos += len(line)
- return line[:-1]
- try:
- chunk = next(self.byte_iterator)
- except StopIteration:
- if self.read_pos < self.buffer.getbuffer().nbytes:
- continue
- raise
- if "PayloadPart" not in chunk:
- # Unknown Event Type
- continue
- self.buffer.seek(0, io.SEEK_END)
- self.buffer.write(chunk["PayloadPart"]["Bytes"])
-
-
-class ContentHandlerBase(Generic[INPUT_TYPE, OUTPUT_TYPE]):
- """Handler class to transform input from LLM to a
- format that SageMaker endpoint expects.
-
- Similarly, the class handles transforming output from the
- SageMaker endpoint to a format that LLM class expects.
- """
-
- """
- Example:
- .. code-block:: python
-
- class ContentHandler(ContentHandlerBase):
- content_type = "application/json"
- accepts = "application/json"
-
- def transform_input(self, prompt: str, model_kwargs: Dict) -> bytes:
- input_str = json.dumps({prompt: prompt, **model_kwargs})
- return input_str.encode('utf-8')
-
- def transform_output(self, output: bytes) -> str:
- response_json = json.loads(output.read().decode("utf-8"))
- return response_json[0]["generated_text"]
- """
-
- content_type: Optional[str] = "text/plain"
- """The MIME type of the input data passed to endpoint"""
-
- accepts: Optional[str] = "text/plain"
- """The MIME type of the response data returned from endpoint"""
-
- @abstractmethod
- def transform_input(self, prompt: INPUT_TYPE, model_kwargs: Dict) -> bytes:
- """Transforms the input to a format that model can accept
- as the request Body. Should return bytes or seekable file
- like object in the format specified in the content_type
- request header.
- """
-
- @abstractmethod
- def transform_output(self, output: bytes) -> OUTPUT_TYPE:
- """Transforms the output from the model to string that
- the LLM class expects.
- """
-
-
-class LLMContentHandler(ContentHandlerBase[str, str]):
- """Content handler for LLM class."""
-
-
-@deprecated(
- since="0.3.16",
- removal="1.0",
- alternative_import="langchain_aws.llms.SagemakerEndpoint",
-)
-class SagemakerEndpoint(LLM):
- """Sagemaker Inference Endpoint models.
-
- To use, you must supply the endpoint name from your deployed
- Sagemaker model & the region where it is deployed.
-
- To authenticate, the AWS client uses the following methods to
- automatically load credentials:
- https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html
-
- If a specific credential profile should be used, you must pass
- the name of the profile from the ~/.aws/credentials file that is to be used.
-
- Make sure the credentials / roles used have the required policies to
- access the Sagemaker endpoint.
- See: https://docs.aws.amazon.com/IAM/latest/UserGuide/access_policies.html
- """
-
- """
- Args:
-
- region_name: The aws region e.g., `us-west-2`.
- Fallsback to AWS_DEFAULT_REGION env variable
- or region specified in ~/.aws/config.
-
- credentials_profile_name: The name of the profile in the ~/.aws/credentials
- or ~/.aws/config files, which has either access keys or role information
- specified. If not specified, the default credential profile or, if on an
- EC2 instance, credentials from IMDS will be used.
-
- client: boto3 client for Sagemaker Endpoint
-
- content_handler: Implementation for model specific LLMContentHandler
-
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import SagemakerEndpoint
- endpoint_name = (
- "my-endpoint-name"
- )
- region_name = (
- "us-west-2"
- )
- credentials_profile_name = (
- "default"
- )
- se = SagemakerEndpoint(
- endpoint_name=endpoint_name,
- region_name=region_name,
- credentials_profile_name=credentials_profile_name
- )
-
- #Use with boto3 client
- client = boto3.client(
- "sagemaker-runtime",
- region_name=region_name
- )
-
- se = SagemakerEndpoint(
- endpoint_name=endpoint_name,
- client=client
- )
-
- """
- client: Any = None
- """Boto3 client for sagemaker runtime"""
-
- endpoint_name: str = ""
- """The name of the endpoint from the deployed Sagemaker model.
- Must be unique within an AWS Region."""
-
- region_name: str = ""
- """The aws region where the Sagemaker model is deployed, eg. `us-west-2`."""
-
- credentials_profile_name: Optional[str] = None
- """The name of the profile in the ~/.aws/credentials or ~/.aws/config files, which
- has either access keys or role information specified.
- If not specified, the default credential profile or, if on an EC2 instance,
- credentials from IMDS will be used.
- See: https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html
- """
-
- content_handler: LLMContentHandler
- """The content handler class that provides an input and
- output transform functions to handle formats between LLM
- and the endpoint.
- """
-
- streaming: bool = False
- """Whether to stream the results."""
-
- """
- Example:
- .. code-block:: python
-
- from langchain_community.llms.sagemaker_endpoint import LLMContentHandler
-
- class ContentHandler(LLMContentHandler):
- content_type = "application/json"
- accepts = "application/json"
-
- def transform_input(self, prompt: str, model_kwargs: Dict) -> bytes:
- input_str = json.dumps({prompt: prompt, **model_kwargs})
- return input_str.encode('utf-8')
-
- def transform_output(self, output: bytes) -> str:
- response_json = json.loads(output.read().decode("utf-8"))
- return response_json[0]["generated_text"]
- """
-
- model_kwargs: Optional[Dict] = None
- """Keyword arguments to pass to the model."""
-
- endpoint_kwargs: Optional[Dict] = None
- """Optional attributes passed to the invoke_endpoint
- function. See `boto3`_. docs for more info.
- .. _boto3:
- """
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Dont do anything if client provided externally"""
- if values.get("client") is not None:
- return values
-
- """Validate that AWS credentials to and python package exists in environment."""
- try:
- import boto3
-
- try:
- if values["credentials_profile_name"] is not None:
- session = boto3.Session(
- profile_name=values["credentials_profile_name"]
- )
- else:
- # use default credentials
- session = boto3.Session()
-
- values["client"] = session.client(
- "sagemaker-runtime", region_name=values["region_name"]
- )
-
- except Exception as e:
- raise ValueError(
- "Could not load credentials to authenticate with AWS client. "
- "Please check that credentials in the specified "
- "profile name are valid."
- ) from e
-
- except ImportError:
- raise ImportError(
- "Could not import boto3 python package. "
- "Please install it with `pip install boto3`."
- )
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- **{"endpoint_name": self.endpoint_name},
- **{"model_kwargs": _model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "sagemaker_endpoint"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Sagemaker inference endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = se("Tell me a joke.")
- """
- _model_kwargs = self.model_kwargs or {}
- _model_kwargs = {**_model_kwargs, **kwargs}
- _endpoint_kwargs = self.endpoint_kwargs or {}
-
- body = self.content_handler.transform_input(prompt, _model_kwargs)
- content_type = self.content_handler.content_type
- accepts = self.content_handler.accepts
-
- if self.streaming and run_manager:
- try:
- resp = self.client.invoke_endpoint_with_response_stream(
- EndpointName=self.endpoint_name,
- Body=body,
- ContentType=self.content_handler.content_type,
- **_endpoint_kwargs,
- )
- iterator = LineIterator(resp["Body"])
- current_completion: str = ""
- for line in iterator:
- resp = json.loads(line)
- resp_output = resp.get("outputs")[0]
- if stop is not None:
- # Uses same approach as below
- resp_output = enforce_stop_tokens(resp_output, stop)
- current_completion += resp_output
- run_manager.on_llm_new_token(resp_output)
- return current_completion
- except Exception as e:
- raise ValueError(f"Error raised by streaming inference endpoint: {e}")
- else:
- try:
- response = self.client.invoke_endpoint(
- EndpointName=self.endpoint_name,
- Body=body,
- ContentType=content_type,
- Accept=accepts,
- **_endpoint_kwargs,
- )
- except Exception as e:
- raise ValueError(f"Error raised by inference endpoint: {e}")
-
- text = self.content_handler.transform_output(response["Body"])
- if stop is not None:
- # This is a bit hacky, but I can't figure out a better way to enforce
- # stop tokens when making calls to the sagemaker endpoint.
- text = enforce_stop_tokens(text, stop)
-
- return text
diff --git a/libs/community/langchain_community/llms/sambanova.py b/libs/community/langchain_community/llms/sambanova.py
deleted file mode 100644
index 18f2810262..0000000000
--- a/libs/community/langchain_community/llms/sambanova.py
+++ /dev/null
@@ -1,866 +0,0 @@
-import json
-from typing import Any, Dict, Iterator, List, Optional, Tuple, Union
-
-import requests
-from langchain_core.callbacks.manager import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env
-from pydantic import Field, SecretStr
-from requests import Response
-
-
-class SambaStudio(LLM):
- """
- SambaStudio large language models.
-
- Setup:
- To use, you should have the environment variables
- ``SAMBASTUDIO_URL`` set with your SambaStudio environment URL.
- ``SAMBASTUDIO_API_KEY`` set with your SambaStudio endpoint API key.
- https://sambanova.ai/products/enterprise-ai-platform-sambanova-suite
- read extra documentation in https://docs.sambanova.ai/sambastudio/latest/index.html
- Example:
- .. code-block:: python
- from langchain_community.llms.sambanova import SambaStudio
- SambaStudio(
- sambastudio_url="your-SambaStudio-environment-URL",
- sambastudio_api_key="your-SambaStudio-API-key,
- model_kwargs={
- "model" : model or expert name (set for Bundle endpoints),
- "max_tokens" : max number of tokens to generate,
- "temperature" : model temperature,
- "top_p" : model top p,
- "top_k" : model top k,
- "do_sample" : wether to do sample
- "process_prompt": wether to process prompt
- (set for Bundle generic v1 and v2 endpoints)
- },
- )
- Key init args — completion params:
- model: str
- The name of the model to use, e.g., Meta-Llama-3-70B-Instruct-4096
- (set for Bundle endpoints).
- streaming: bool
- Whether to use streaming handler when using non streaming methods
- model_kwargs: dict
- Extra Key word arguments to pass to the model:
- max_tokens: int
- max tokens to generate
- temperature: float
- model temperature
- top_p: float
- model top p
- top_k: int
- model top k
- do_sample: bool
- wether to do sample
- process_prompt:
- wether to process prompt
- (set for Bundle generic v1 and v2 endpoints)
- Key init args — client params:
- sambastudio_url: str
- SambaStudio endpoint Url
- sambastudio_api_key: str
- SambaStudio endpoint api key
-
- Instantiate:
- .. code-block:: python
-
- from langchain_community.llms import SambaStudio
-
- llm = SambaStudio=(
- sambastudio_url = set with your SambaStudio deployed endpoint URL,
- sambastudio_api_key = set with your SambaStudio deployed endpoint Key,
- model_kwargs = {
- "model" : model or expert name (set for Bundle endpoints),
- "max_tokens" : max number of tokens to generate,
- "temperature" : model temperature,
- "top_p" : model top p,
- "top_k" : model top k,
- "do_sample" : wether to do sample
- "process_prompt" : wether to process prompt
- (set for Bundle generic v1 and v2 endpoints)
- }
- )
-
- Invoke:
- .. code-block:: python
- prompt = "tell me a joke"
- response = llm.invoke(prompt)
-
- Stream:
- .. code-block:: python
-
- for chunk in llm.stream(prompt):
- print(chunk, end="", flush=True)
-
- Async:
- .. code-block:: python
-
- response = llm.ainvoke(prompt)
- await response
-
- """
-
- sambastudio_url: str = Field(default="")
- """SambaStudio Url"""
-
- sambastudio_api_key: SecretStr = Field(default=SecretStr(""))
- """SambaStudio api key"""
-
- base_url: str = Field(default="", exclude=True)
- """SambaStudio non streaming URL"""
-
- streaming_url: str = Field(default="", exclude=True)
- """SambaStudio streaming URL"""
-
- streaming: bool = Field(default=False)
- """Whether to use streaming handler when using non streaming methods"""
-
- model_kwargs: Optional[Dict[str, Any]] = None
- """Key word arguments to pass to the model."""
-
- class Config:
- populate_by_name = True
-
- @classmethod
- def is_lc_serializable(cls) -> bool:
- """Return whether this model can be serialized by Langchain."""
- return True
-
- @property
- def lc_secrets(self) -> Dict[str, str]:
- return {
- "sambastudio_url": "sambastudio_url",
- "sambastudio_api_key": "sambastudio_api_key",
- }
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Return a dictionary of identifying parameters.
-
- This information is used by the LangChain callback system, which
- is used for tracing purposes make it possible to monitor LLMs.
- """
- return {"streaming": self.streaming, **{"model_kwargs": self.model_kwargs}}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "sambastudio-llm"
-
- def __init__(self, **kwargs: Any) -> None:
- """init and validate environment variables"""
- kwargs["sambastudio_url"] = get_from_dict_or_env(
- kwargs, "sambastudio_url", "SAMBASTUDIO_URL"
- )
-
- kwargs["sambastudio_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(kwargs, "sambastudio_api_key", "SAMBASTUDIO_API_KEY")
- )
- kwargs["base_url"], kwargs["streaming_url"] = self._get_sambastudio_urls(
- kwargs["sambastudio_url"]
- )
- super().__init__(**kwargs)
-
- def _get_sambastudio_urls(self, url: str) -> Tuple[str, str]:
- """
- Get streaming and non streaming URLs from the given URL
-
- Args:
- url: string with sambastudio base or streaming endpoint url
-
- Returns:
- base_url: string with url to do non streaming calls
- streaming_url: string with url to do streaming calls
- """
- if "chat/completions" in url:
- base_url = url
- stream_url = url
- else:
- if "stream" in url:
- base_url = url.replace("stream/", "")
- stream_url = url
- else:
- base_url = url
- if "generic" in url:
- stream_url = "generic/stream".join(url.split("generic"))
- else:
- raise ValueError("Unsupported URL")
- return base_url, stream_url
-
- def _get_tuning_params(self, stop: Optional[List[str]] = None) -> Dict[str, Any]:
- """
- Get the tuning parameters to use when calling the LLM.
-
- Args:
- stop: Stop words to use when generating. Model output is cut off at the
- first occurrence of any of the stop substrings.
-
- Returns:
- The tuning parameters in the format required by api to use
- """
- if stop is None:
- stop = []
-
- # get the parameters to use when calling the LLM.
- _model_kwargs = self.model_kwargs or {}
-
- # handle the case where stop sequences are send in the invocation
- # and stop sequences has been also set in the model parameters
- _stop_sequences = _model_kwargs.get("stop_sequences", []) + stop
- if len(_stop_sequences) > 0:
- _model_kwargs["stop_sequences"] = _stop_sequences
-
- # set the parameters structure depending of the API
- if "chat/completions" in self.sambastudio_url:
- if "select_expert" in _model_kwargs.keys():
- _model_kwargs["model"] = _model_kwargs.pop("select_expert")
- if "max_tokens_to_generate" in _model_kwargs.keys():
- _model_kwargs["max_tokens"] = _model_kwargs.pop(
- "max_tokens_to_generate"
- )
- if "process_prompt" in _model_kwargs.keys():
- _model_kwargs.pop("process_prompt")
- tuning_params = _model_kwargs
-
- elif "api/v2/predict/generic" in self.sambastudio_url:
- if "model" in _model_kwargs.keys():
- _model_kwargs["select_expert"] = _model_kwargs.pop("model")
- if "max_tokens" in _model_kwargs.keys():
- _model_kwargs["max_tokens_to_generate"] = _model_kwargs.pop(
- "max_tokens"
- )
- tuning_params = _model_kwargs
-
- elif "api/predict/generic" in self.sambastudio_url:
- if "model" in _model_kwargs.keys():
- _model_kwargs["select_expert"] = _model_kwargs.pop("model")
- if "max_tokens" in _model_kwargs.keys():
- _model_kwargs["max_tokens_to_generate"] = _model_kwargs.pop(
- "max_tokens"
- )
-
- tuning_params = {
- k: {"type": type(v).__name__, "value": str(v)}
- for k, v in (_model_kwargs.items())
- }
-
- else:
- raise ValueError(
- f"Unsupported URL{self.sambastudio_url}"
- "only openai, generic v1 and generic v2 APIs are supported"
- )
-
- return tuning_params
-
- def _handle_request(
- self,
- prompt: Union[List[str], str],
- stop: Optional[List[str]] = None,
- streaming: Optional[bool] = False,
- ) -> Response:
- """
- Performs a post request to the LLM API.
-
- Args:
- prompt: The prompt to pass into the model
- stop: list of stop tokens
- streaming: wether to do a streaming call
-
- Returns:
- A request Response object
- """
-
- if isinstance(prompt, str):
- prompt = [prompt]
-
- params = self._get_tuning_params(stop)
-
- # create request payload for openAI v1 API
- if "chat/completions" in self.sambastudio_url:
- messages_dict = [{"role": "user", "content": prompt[0]}]
- data = {"messages": messages_dict, "stream": streaming, **params}
- data = {key: value for key, value in data.items() if value is not None}
- headers = {
- "Authorization": f"Bearer "
- f"{self.sambastudio_api_key.get_secret_value()}",
- "Content-Type": "application/json",
- }
-
- # create request payload for generic v1 API
- elif "api/v2/predict/generic" in self.sambastudio_url:
- if params.get("process_prompt", False):
- prompt = json.dumps(
- {
- "conversation_id": "sambaverse-conversation-id",
- "messages": [
- {"message_id": None, "role": "user", "content": prompt[0]}
- ],
- }
- )
- else:
- prompt = prompt[0]
- items = [{"id": "item0", "value": prompt}]
- params = {key: value for key, value in params.items() if value is not None}
- data = {"items": items, "params": params}
- headers = {"key": self.sambastudio_api_key.get_secret_value()}
-
- # create request payload for generic v1 API
- elif "api/predict/generic" in self.sambastudio_url:
- if params.get("process_prompt", False):
- if params["process_prompt"].get("value") == "True":
- prompt = json.dumps(
- {
- "conversation_id": "sambaverse-conversation-id",
- "messages": [
- {
- "message_id": None,
- "role": "user",
- "content": prompt[0],
- }
- ],
- }
- )
- else:
- prompt = prompt[0]
- else:
- prompt = prompt[0]
- if streaming:
- data = {"instance": prompt, "params": params}
- else:
- data = {"instances": [prompt], "params": params}
- headers = {"key": self.sambastudio_api_key.get_secret_value()}
-
- else:
- raise ValueError(
- f"Unsupported URL{self.sambastudio_url}"
- "only openai, generic v1 and generic v2 APIs are supported"
- )
-
- # make the request to SambaStudio API
- http_session = requests.Session()
- if streaming:
- response = http_session.post(
- self.streaming_url, headers=headers, json=data, stream=True
- )
- else:
- response = http_session.post(
- self.base_url, headers=headers, json=data, stream=False
- )
- if response.status_code != 200:
- raise RuntimeError(
- f"Sambanova / complete call failed with status code "
- f"{response.status_code}."
- f"{response.text}."
- )
- return response
-
- def _process_response(self, response: Response) -> str:
- """
- Process a non streaming response from the api
-
- Args:
- response: A request Response object
-
- Returns
- completion: a string with model generation
- """
-
- # Extract json payload form response
- try:
- response_dict = response.json()
- except Exception as e:
- raise RuntimeError(
- f"Sambanova /complete call failed couldn't get JSON response {e}"
- f"response: {response.text}"
- )
-
- # process response payload for openai compatible API
- if "chat/completions" in self.sambastudio_url:
- completion = response_dict["choices"][0]["message"]["content"]
- # process response payload for generic v2 API
- elif "api/v2/predict/generic" in self.sambastudio_url:
- completion = response_dict["items"][0]["value"]["completion"]
- # process response payload for generic v1 API
- elif "api/predict/generic" in self.sambastudio_url:
- completion = response_dict["predictions"][0]["completion"]
- else:
- raise ValueError(
- f"Unsupported URL{self.sambastudio_url}"
- "only openai, generic v1 and generic v2 APIs are supported"
- )
- return completion
-
- def _process_stream_response(self, response: Response) -> Iterator[GenerationChunk]:
- """
- Process a streaming response from the api
-
- Args:
- response: An iterable request Response object
-
- Yields:
- GenerationChunk: a GenerationChunk with model partial generation
- """
-
- try:
- import sseclient
- except ImportError:
- raise ImportError(
- "could not import sseclient library"
- "Please install it with `pip install sseclient-py`."
- )
-
- # process response payload for openai compatible API
- if "chat/completions" in self.sambastudio_url:
- client = sseclient.SSEClient(response)
- for event in client.events():
- if event.event == "error_event":
- raise RuntimeError(
- f"Sambanova /complete call failed with status code "
- f"{response.status_code}."
- f"{event.data}."
- )
- try:
- # check if the response is not a final event ("[DONE]")
- if event.data != "[DONE]":
- if isinstance(event.data, str):
- data = json.loads(event.data)
- else:
- raise RuntimeError(
- f"Sambanova /complete call failed with status code "
- f"{response.status_code}."
- f"{event.data}."
- )
- if data.get("error"):
- raise RuntimeError(
- f"Sambanova /complete call failed with status code "
- f"{response.status_code}."
- f"{event.data}."
- )
- if len(data["choices"]) > 0:
- content = data["choices"][0]["delta"]["content"]
- else:
- content = ""
- generated_chunk = GenerationChunk(text=content)
- yield generated_chunk
-
- except Exception as e:
- raise RuntimeError(
- f"Error getting content chunk raw streamed response: {e}"
- f"data: {event.data}"
- )
-
- # process response payload for generic v2 API
- elif "api/v2/predict/generic" in self.sambastudio_url:
- for line in response.iter_lines():
- try:
- data = json.loads(line)
- content = data["result"]["items"][0]["value"]["stream_token"]
- generated_chunk = GenerationChunk(text=content)
- yield generated_chunk
-
- except Exception as e:
- raise RuntimeError(
- f"Error getting content chunk raw streamed response: {e}"
- f"line: {line}"
- )
-
- # process response payload for generic v1 API
- elif "api/predict/generic" in self.sambastudio_url:
- for line in response.iter_lines():
- try:
- data = json.loads(line)
- content = data["result"]["responses"][0]["stream_token"]
- generated_chunk = GenerationChunk(text=content)
- yield generated_chunk
-
- except Exception as e:
- raise RuntimeError(
- f"Error getting content chunk raw streamed response: {e}"
- f"line: {line}"
- )
-
- else:
- raise ValueError(
- f"Unsupported URL{self.sambastudio_url}"
- "only openai, generic v1 and generic v2 APIs are supported"
- )
-
- def _stream(
- self,
- prompt: Union[List[str], str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- """Call out to Sambanova's complete endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: a list of strings on which the model should stop generating.
- run_manager: A run manager with callbacks for the LLM.
- Yields:
- chunk: GenerationChunk with model partial generation
- """
- response = self._handle_request(prompt, stop, streaming=True)
- for chunk in self._process_stream_response(response):
- if run_manager:
- run_manager.on_llm_new_token(chunk.text)
- yield chunk
-
- def _call(
- self,
- prompt: Union[List[str], str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Sambanova's complete endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: a list of strings on which the model should stop generating.
-
- Returns:
- result: string with model generation
- """
- if self.streaming:
- completion = ""
- for chunk in self._stream(
- prompt=prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- completion += chunk.text
-
- return completion
-
- response = self._handle_request(prompt, stop, streaming=False)
- completion = self._process_response(response)
- return completion
-
-
-class SambaNovaCloud(LLM):
- """
- SambaNova Cloud large language models.
-
- Setup:
- To use, you should have the environment variables:
- ``SAMBANOVA_URL`` set with SambaNova Cloud URL.
- defaults to http://cloud.sambanova.ai/
- ``SAMBANOVA_API_KEY`` set with your SambaNova Cloud API Key.
- Example:
- .. code-block:: python
- from langchain_community.llms.sambanova import SambaNovaCloud
- SambaNovaCloud(
- sambanova_api_key="your-SambaNovaCloud-API-key,
- model = model name,
- max_tokens = max number of tokens to generate,
- temperature = model temperature,
- top_p = model top p,
- top_k = model top k
- )
- Key init args — completion params:
- model: str
- The name of the model to use, e.g., Meta-Llama-3-70B-Instruct-4096
- (set for CoE endpoints).
- streaming: bool
- Whether to use streaming handler when using non streaming methods
- max_tokens: int
- max tokens to generate
- temperature: float
- model temperature
- top_p: float
- model top p
- top_k: int
- model top k
-
- Key init args — client params:
- sambanova_url: str
- SambaNovaCloud Url defaults to http://cloud.sambanova.ai/
- sambanova_api_key: str
- SambaNovaCloud api key
- Instantiate:
- .. code-block:: python
- from langchain_community.llms.sambanova import SambaNovaCloud
- SambaNovaCloud(
- sambanova_api_key="your-SambaNovaCloud-API-key,
- model = model name,
- max_tokens = max number of tokens to generate,
- temperature = model temperature,
- top_p = model top p,
- top_k = model top k
- )
- Invoke:
- .. code-block:: python
- prompt = "tell me a joke"
- response = llm.invoke(prompt)
- Stream:
- .. code-block:: python
- for chunk in llm.stream(prompt):
- print(chunk, end="", flush=True)
- Async:
- .. code-block:: python
- response = llm.ainvoke(prompt)
- await response
- """
-
- sambanova_url: str = Field(default="")
- """SambaNova Cloud Url"""
-
- sambanova_api_key: SecretStr = Field(default=SecretStr(""))
- """SambaNova Cloud api key"""
-
- model: str = Field(default="Meta-Llama-3.1-8B-Instruct")
- """The name of the model"""
-
- streaming: bool = Field(default=False)
- """Whether to use streaming handler when using non streaming methods"""
-
- max_tokens: int = Field(default=1024)
- """max tokens to generate"""
-
- temperature: float = Field(default=0.7)
- """model temperature"""
-
- top_p: Optional[float] = Field(default=None)
- """model top p"""
-
- top_k: Optional[int] = Field(default=None)
- """model top k"""
-
- stream_options: dict = Field(default={"include_usage": True})
- """stream options, include usage to get generation metrics"""
-
- class Config:
- populate_by_name = True
-
- @classmethod
- def is_lc_serializable(cls) -> bool:
- """Return whether this model can be serialized by Langchain."""
- return False
-
- @property
- def lc_secrets(self) -> Dict[str, str]:
- return {"sambanova_api_key": "sambanova_api_key"}
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Return a dictionary of identifying parameters.
-
- This information is used by the LangChain callback system, which
- is used for tracing purposes make it possible to monitor LLMs.
- """
- return {
- "model": self.model,
- "streaming": self.streaming,
- "max_tokens": self.max_tokens,
- "temperature": self.temperature,
- "top_p": self.top_p,
- "top_k": self.top_k,
- "stream_options": self.stream_options,
- }
-
- @property
- def _llm_type(self) -> str:
- """Get the type of language model used by this chat model."""
- return "sambanovacloud-llm"
-
- def __init__(self, **kwargs: Any) -> None:
- """init and validate environment variables"""
- kwargs["sambanova_url"] = get_from_dict_or_env(
- kwargs,
- "sambanova_url",
- "SAMBANOVA_URL",
- default="https://api.sambanova.ai/v1/chat/completions",
- )
- kwargs["sambanova_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(kwargs, "sambanova_api_key", "SAMBANOVA_API_KEY")
- )
- super().__init__(**kwargs)
-
- def _handle_request(
- self,
- prompt: Union[List[str], str],
- stop: Optional[List[str]] = None,
- streaming: Optional[bool] = False,
- ) -> Response:
- """
- Performs a post request to the LLM API.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: list of stop tokens
-
- Returns:
- A request Response object
- """
- if isinstance(prompt, str):
- prompt = [prompt]
-
- messages_dict = [{"role": "user", "content": prompt[0]}]
- data = {
- "messages": messages_dict,
- "stream": streaming,
- "max_tokens": self.max_tokens,
- "stop": stop,
- "model": self.model,
- "temperature": self.temperature,
- "top_p": self.top_p,
- "top_k": self.top_k,
- }
- data = {key: value for key, value in data.items() if value is not None}
- headers = {
- "Authorization": f"Bearer {self.sambanova_api_key.get_secret_value()}",
- "Content-Type": "application/json",
- }
-
- http_session = requests.Session()
- if streaming:
- response = http_session.post(
- self.sambanova_url, headers=headers, json=data, stream=True
- )
- else:
- response = http_session.post(
- self.sambanova_url, headers=headers, json=data, stream=False
- )
-
- if response.status_code != 200:
- raise RuntimeError(
- f"Sambanova / complete call failed with status code "
- f"{response.status_code}."
- f"{response.text}."
- )
- return response
-
- def _process_response(self, response: Response) -> str:
- """
- Process a non streaming response from the api
-
- Args:
- response: A request Response object
-
- Returns
- completion: a string with model generation
- """
-
- # Extract json payload form response
- try:
- response_dict = response.json()
- except Exception as e:
- raise RuntimeError(
- f"Sambanova /complete call failed couldn't get JSON response {e}"
- f"response: {response.text}"
- )
-
- completion = response_dict["choices"][0]["message"]["content"]
-
- return completion
-
- def _process_stream_response(self, response: Response) -> Iterator[GenerationChunk]:
- """
- Process a streaming response from the api
-
- Args:
- response: An iterable request Response object
-
- Yields:
- GenerationChunk: a GenerationChunk with model partial generation
- """
-
- try:
- import sseclient
- except ImportError:
- raise ImportError(
- "could not import sseclient library"
- "Please install it with `pip install sseclient-py`."
- )
-
- client = sseclient.SSEClient(response)
- for event in client.events():
- if event.event == "error_event":
- raise RuntimeError(
- f"Sambanova /complete call failed with status code "
- f"{response.status_code}."
- f"{event.data}."
- )
- try:
- # check if the response is not a final event ("[DONE]")
- if event.data != "[DONE]":
- if isinstance(event.data, str):
- data = json.loads(event.data)
- else:
- raise RuntimeError(
- f"Sambanova /complete call failed with status code "
- f"{response.status_code}."
- f"{event.data}."
- )
- if data.get("error"):
- raise RuntimeError(
- f"Sambanova /complete call failed with status code "
- f"{response.status_code}."
- f"{event.data}."
- )
- if len(data["choices"]) > 0:
- content = data["choices"][0]["delta"]["content"]
- else:
- content = ""
- generated_chunk = GenerationChunk(text=content)
- yield generated_chunk
-
- except Exception as e:
- raise RuntimeError(
- f"Error getting content chunk raw streamed response: {e}"
- f"data: {event.data}"
- )
-
- def _call(
- self,
- prompt: Union[List[str], str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to SambaNovaCloud complete endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
- """
- if self.streaming:
- completion = ""
- for chunk in self._stream(
- prompt=prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- completion += chunk.text
-
- return completion
-
- response = self._handle_request(prompt, stop, streaming=False)
- completion = self._process_response(response)
- return completion
-
- def _stream(
- self,
- prompt: Union[List[str], str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- """Call out to SambaNovaCloud complete endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
- """
- response = self._handle_request(prompt, stop, streaming=True)
- for chunk in self._process_stream_response(response):
- if run_manager:
- run_manager.on_llm_new_token(chunk.text)
- yield chunk
diff --git a/libs/community/langchain_community/llms/self_hosted.py b/libs/community/langchain_community/llms/self_hosted.py
deleted file mode 100644
index 7009868578..0000000000
--- a/libs/community/langchain_community/llms/self_hosted.py
+++ /dev/null
@@ -1,236 +0,0 @@
-import importlib.util
-import logging
-import pickle
-from typing import Any, Callable, List, Mapping, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from pydantic import ConfigDict
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-def _generate_text(
- pipeline: Any,
- prompt: str,
- *args: Any,
- stop: Optional[List[str]] = None,
- **kwargs: Any,
-) -> str:
- """Inference function to send to the remote hardware.
-
- Accepts a pipeline callable (or, more likely,
- a key pointing to the model on the cluster's object store)
- and returns text predictions for each document
- in the batch.
- """
- text = pipeline(prompt, *args, **kwargs)
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
- return text
-
-
-def _send_pipeline_to_device(pipeline: Any, device: int) -> Any:
- """Send a pipeline to a device on the cluster."""
- if isinstance(pipeline, str):
- with open(pipeline, "rb") as f:
- # This code path can only be triggered if the user
- # passed allow_dangerous_deserialization=True
- pipeline = pickle.load(f) # ignore[pickle]: explicit-opt-in
-
- if importlib.util.find_spec("torch") is not None:
- import torch
-
- cuda_device_count = torch.cuda.device_count()
- if device < -1 or (device >= cuda_device_count):
- raise ValueError(
- f"Got device=={device}, "
- f"device is required to be within [-1, {cuda_device_count})"
- )
- if device < 0 and cuda_device_count > 0:
- logger.warning(
- "Device has %d GPUs available. "
- "Provide device={deviceId} to `from_model_id` to use available"
- "GPUs for execution. deviceId is -1 for CPU and "
- "can be a positive integer associated with CUDA device id.",
- cuda_device_count,
- )
-
- pipeline.device = torch.device(device)
- pipeline.model = pipeline.model.to(pipeline.device)
- return pipeline
-
-
-class SelfHostedPipeline(LLM):
- """Model inference on self-hosted remote hardware.
-
- Supported hardware includes auto-launched instances on AWS, GCP, Azure,
- and Lambda, as well as servers specified
- by IP address and SSH credentials (such as on-prem, or another
- cloud like Paperspace, Coreweave, etc.).
-
- To use, you should have the ``runhouse`` python package installed.
-
- Example for custom pipeline and inference functions:
- .. code-block:: python
-
- from langchain_community.llms import SelfHostedPipeline
- from transformers import AutoModelForCausalLM, AutoTokenizer, pipeline
- import runhouse as rh
-
- def load_pipeline():
- tokenizer = AutoTokenizer.from_pretrained("gpt2")
- model = AutoModelForCausalLM.from_pretrained("gpt2")
- return pipeline(
- "text-generation", model=model, tokenizer=tokenizer,
- max_new_tokens=10
- )
- def inference_fn(pipeline, prompt, stop = None):
- return pipeline(prompt)[0]["generated_text"]
-
- gpu = rh.cluster(name="rh-a10x", instance_type="A100:1")
- llm = SelfHostedPipeline(
- model_load_fn=load_pipeline,
- hardware=gpu,
- model_reqs=model_reqs, inference_fn=inference_fn
- )
- Example for <2GB model (can be serialized and sent directly to the server):
- .. code-block:: python
-
- from langchain_community.llms import SelfHostedPipeline
- import runhouse as rh
- gpu = rh.cluster(name="rh-a10x", instance_type="A100:1")
- my_model = ...
- llm = SelfHostedPipeline.from_pipeline(
- pipeline=my_model,
- hardware=gpu,
- model_reqs=["./", "torch", "transformers"],
- )
- Example passing model path for larger models:
- .. code-block:: python
-
- from langchain_community.llms import SelfHostedPipeline
- import runhouse as rh
- import pickle
- from transformers import pipeline
-
- generator = pipeline(model="gpt2")
- rh.blob(pickle.dumps(generator), path="models/pipeline.pkl"
- ).save().to(gpu, path="models")
- llm = SelfHostedPipeline.from_pipeline(
- pipeline="models/pipeline.pkl",
- hardware=gpu,
- model_reqs=["./", "torch", "transformers"],
- )
- """
-
- pipeline_ref: Any = None #: :meta private:
- client: Any = None #: :meta private:
- inference_fn: Callable = _generate_text #: :meta private:
- """Inference function to send to the remote hardware."""
- hardware: Any = None
- """Remote hardware to send the inference function to."""
- model_load_fn: Callable
- """Function to load the model remotely on the server."""
- load_fn_kwargs: Optional[dict] = None
- """Keyword arguments to pass to the model load function."""
- model_reqs: List[str] = ["./", "torch"]
- """Requirements to install on hardware to inference the model."""
-
- allow_dangerous_deserialization: bool = False
- """Allow deserialization using pickle which can be dangerous if
- loading compromised data.
- """
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- def __init__(self, **kwargs: Any):
- """Init the pipeline with an auxiliary function.
-
- The load function must be in global scope to be imported
- and run on the server, i.e. in a module and not a REPL or closure.
- Then, initialize the remote inference function.
- """
- if not kwargs.get("allow_dangerous_deserialization"):
- raise ValueError(
- "SelfHostedPipeline relies on the pickle module. "
- "You will need to set allow_dangerous_deserialization=True "
- "if you want to opt-in to allow deserialization of data using pickle."
- "Data can be compromised by a malicious actor if "
- "not handled properly to include "
- "a malicious payload that when deserialized with "
- "pickle can execute arbitrary code. "
- )
- super().__init__(**kwargs)
- try:
- import runhouse as rh
-
- except ImportError:
- raise ImportError(
- "Could not import runhouse python package. "
- "Please install it with `pip install runhouse`."
- )
-
- remote_load_fn = rh.function(fn=self.model_load_fn).to(
- self.hardware, reqs=self.model_reqs
- )
- _load_fn_kwargs = self.load_fn_kwargs or {}
- self.pipeline_ref = remote_load_fn.remote(**_load_fn_kwargs)
-
- self.client = rh.function(fn=self.inference_fn).to(
- self.hardware, reqs=self.model_reqs
- )
-
- @classmethod
- def from_pipeline(
- cls,
- pipeline: Any,
- hardware: Any,
- model_reqs: Optional[List[str]] = None,
- device: int = 0,
- **kwargs: Any,
- ) -> LLM:
- """Init the SelfHostedPipeline from a pipeline object or string."""
- if not isinstance(pipeline, str):
- logger.warning(
- "Serializing pipeline to send to remote hardware. "
- "Note, it can be quite slow"
- "to serialize and send large models with each execution. "
- "Consider sending the pipeline"
- "to the cluster and passing the path to the pipeline instead."
- )
-
- load_fn_kwargs = {"pipeline": pipeline, "device": device}
- return cls(
- load_fn_kwargs=load_fn_kwargs,
- model_load_fn=_send_pipeline_to_device,
- hardware=hardware,
- model_reqs=["transformers", "torch"] + (model_reqs or []),
- **kwargs,
- )
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- **{"hardware": self.hardware},
- }
-
- @property
- def _llm_type(self) -> str:
- return "self_hosted_llm"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- return self.client(
- pipeline=self.pipeline_ref, prompt=prompt, stop=stop, **kwargs
- )
diff --git a/libs/community/langchain_community/llms/self_hosted_hugging_face.py b/libs/community/langchain_community/llms/self_hosted_hugging_face.py
deleted file mode 100644
index e43ca9e312..0000000000
--- a/libs/community/langchain_community/llms/self_hosted_hugging_face.py
+++ /dev/null
@@ -1,211 +0,0 @@
-import importlib.util
-import logging
-from typing import Any, Callable, List, Mapping, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from pydantic import ConfigDict
-
-from langchain_community.llms.self_hosted import SelfHostedPipeline
-from langchain_community.llms.utils import enforce_stop_tokens
-
-DEFAULT_MODEL_ID = "gpt2"
-DEFAULT_TASK = "text-generation"
-VALID_TASKS = ("text2text-generation", "text-generation", "summarization")
-
-logger = logging.getLogger(__name__)
-
-
-def _generate_text(
- pipeline: Any,
- prompt: str,
- *args: Any,
- stop: Optional[List[str]] = None,
- **kwargs: Any,
-) -> str:
- """Inference function to send to the remote hardware.
-
- Accepts a Hugging Face pipeline (or more likely,
- a key pointing to such a pipeline on the cluster's object store)
- and returns generated text.
- """
- response = pipeline(prompt, *args, **kwargs)
- if pipeline.task == "text-generation":
- # Text generation return includes the starter text.
- text = response[0]["generated_text"][len(prompt) :]
- elif pipeline.task == "text2text-generation":
- text = response[0]["generated_text"]
- elif pipeline.task == "summarization":
- text = response[0]["summary_text"]
- else:
- raise ValueError(
- f"Got invalid task {pipeline.task}, "
- f"currently only {VALID_TASKS} are supported"
- )
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
- return text
-
-
-def _load_transformer(
- model_id: str = DEFAULT_MODEL_ID,
- task: str = DEFAULT_TASK,
- device: int = 0,
- model_kwargs: Optional[dict] = None,
-) -> Any:
- """Inference function to send to the remote hardware.
-
- Accepts a huggingface model_id and returns a pipeline for the task.
- """
- from transformers import AutoModelForCausalLM, AutoModelForSeq2SeqLM, AutoTokenizer
- from transformers import pipeline as hf_pipeline
-
- _model_kwargs = model_kwargs or {}
- tokenizer = AutoTokenizer.from_pretrained(model_id, **_model_kwargs)
-
- try:
- if task == "text-generation":
- model = AutoModelForCausalLM.from_pretrained(model_id, **_model_kwargs)
- elif task in ("text2text-generation", "summarization"):
- model = AutoModelForSeq2SeqLM.from_pretrained(model_id, **_model_kwargs)
- else:
- raise ValueError(
- f"Got invalid task {task}, currently only {VALID_TASKS} are supported"
- )
- except ImportError as e:
- raise ImportError(
- f"Could not load the {task} model due to missing dependencies."
- ) from e
-
- if importlib.util.find_spec("torch") is not None:
- import torch
-
- cuda_device_count = torch.cuda.device_count()
- if device < -1 or (device >= cuda_device_count):
- raise ValueError(
- f"Got device=={device}, "
- f"device is required to be within [-1, {cuda_device_count})"
- )
- if device < 0 and cuda_device_count > 0:
- logger.warning(
- "Device has %d GPUs available. "
- "Provide device={deviceId} to `from_model_id` to use available"
- "GPUs for execution. deviceId is -1 for CPU and "
- "can be a positive integer associated with CUDA device id.",
- cuda_device_count,
- )
-
- pipeline = hf_pipeline(
- task=task,
- model=model,
- tokenizer=tokenizer,
- device=device,
- model_kwargs=_model_kwargs,
- )
- if pipeline.task not in VALID_TASKS:
- raise ValueError(
- f"Got invalid task {pipeline.task}, "
- f"currently only {VALID_TASKS} are supported"
- )
- return pipeline
-
-
-class SelfHostedHuggingFaceLLM(SelfHostedPipeline):
- """HuggingFace Pipeline API to run on self-hosted remote hardware.
-
- Supported hardware includes auto-launched instances on AWS, GCP, Azure,
- and Lambda, as well as servers specified
- by IP address and SSH credentials (such as on-prem, or another cloud
- like Paperspace, Coreweave, etc.).
-
- To use, you should have the ``runhouse`` python package installed.
-
- Only supports `text-generation`, `text2text-generation` and `summarization` for now.
-
- Example using from_model_id:
- .. code-block:: python
-
- from langchain_community.llms import SelfHostedHuggingFaceLLM
- import runhouse as rh
- gpu = rh.cluster(name="rh-a10x", instance_type="A100:1")
- hf = SelfHostedHuggingFaceLLM(
- model_id="google/flan-t5-large", task="text2text-generation",
- hardware=gpu
- )
- Example passing fn that generates a pipeline (bc the pipeline is not serializable):
- .. code-block:: python
-
- from langchain_community.llms import SelfHostedHuggingFaceLLM
- from transformers import AutoModelForCausalLM, AutoTokenizer, pipeline
- import runhouse as rh
-
- def get_pipeline():
- model_id = "gpt2"
- tokenizer = AutoTokenizer.from_pretrained(model_id)
- model = AutoModelForCausalLM.from_pretrained(model_id)
- pipe = pipeline(
- "text-generation", model=model, tokenizer=tokenizer
- )
- return pipe
- hf = SelfHostedHuggingFaceLLM(
- model_load_fn=get_pipeline, model_id="gpt2", hardware=gpu)
- """
-
- model_id: str = DEFAULT_MODEL_ID
- """Hugging Face model_id to load the model."""
- task: str = DEFAULT_TASK
- """Hugging Face task ("text-generation", "text2text-generation" or
- "summarization")."""
- device: int = 0
- """Device to use for inference. -1 for CPU, 0 for GPU, 1 for second GPU, etc."""
- model_kwargs: Optional[dict] = None
- """Keyword arguments to pass to the model."""
- hardware: Any = None
- """Remote hardware to send the inference function to."""
- model_reqs: List[str] = ["./", "transformers", "torch"]
- """Requirements to install on hardware to inference the model."""
- model_load_fn: Callable = _load_transformer
- """Function to load the model remotely on the server."""
- inference_fn: Callable = _generate_text #: :meta private:
- """Inference function to send to the remote hardware."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- def __init__(self, **kwargs: Any):
- """Construct the pipeline remotely using an auxiliary function.
-
- The load function needs to be importable to be imported
- and run on the server, i.e. in a module and not a REPL or closure.
- Then, initialize the remote inference function.
- """
- load_fn_kwargs = {
- "model_id": kwargs.get("model_id", DEFAULT_MODEL_ID),
- "task": kwargs.get("task", DEFAULT_TASK),
- "device": kwargs.get("device", 0),
- "model_kwargs": kwargs.get("model_kwargs", None),
- }
- super().__init__(load_fn_kwargs=load_fn_kwargs, **kwargs)
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- **{"model_id": self.model_id},
- **{"model_kwargs": self.model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- return "selfhosted_huggingface_pipeline"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- return self.client(
- pipeline=self.pipeline_ref, prompt=prompt, stop=stop, **kwargs
- )
diff --git a/libs/community/langchain_community/llms/solar.py b/libs/community/langchain_community/llms/solar.py
deleted file mode 100644
index 8bfaecd503..0000000000
--- a/libs/community/langchain_community/llms/solar.py
+++ /dev/null
@@ -1,132 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import (
- BaseModel,
- ConfigDict,
- Field,
- SecretStr,
- model_validator,
-)
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-SOLAR_SERVICE_URL_BASE = "https://api.upstage.ai/v1/solar"
-SOLAR_SERVICE = "https://api.upstage.ai"
-
-
-class _SolarClient(BaseModel):
- """An API client that talks to the Solar server."""
-
- api_key: SecretStr
- """The API key to use for authentication."""
- base_url: str = SOLAR_SERVICE_URL_BASE
-
- def completion(self, request: Any) -> Any:
- headers = {"Authorization": f"Bearer {self.api_key.get_secret_value()}"}
- response = requests.post(
- f"{self.base_url}/chat/completions",
- headers=headers,
- json=request,
- )
- if not response.ok:
- raise ValueError(f"HTTP {response.status_code} error: {response.text}")
- return response.json()["choices"][0]["message"]["content"]
-
-
-class SolarCommon(BaseModel):
- """Common configuration for Solar LLMs."""
-
- _client: _SolarClient
- base_url: str = SOLAR_SERVICE_URL_BASE
- solar_api_key: Optional[SecretStr] = Field(default=None, alias="api_key")
- """Solar API key. Get it here: https://console.upstage.ai/services/solar"""
- model_name: str = Field(default="solar-mini", alias="model")
- """Model name. Available models listed here: https://console.upstage.ai/services/solar"""
- max_tokens: int = Field(default=1024)
- temperature: float = 0.3
-
- model_config = ConfigDict(
- populate_by_name=True,
- arbitrary_types_allowed=True,
- extra="ignore",
- protected_namespaces=(),
- )
-
- @property
- def lc_secrets(self) -> dict:
- return {"solar_api_key": "SOLAR_API_KEY"}
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- return {
- "model": self.model_name,
- "max_tokens": self.max_tokens,
- "temperature": self.temperature,
- }
-
- @property
- def _invocation_params(self) -> Dict[str, Any]:
- return {**{"model": self.model_name}, **self._default_params}
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- api_key = get_from_dict_or_env(values, "solar_api_key", "SOLAR_API_KEY")
- if api_key is None or len(api_key) == 0:
- raise ValueError("SOLAR_API_KEY must be configured")
-
- values["solar_api_key"] = convert_to_secret_str(api_key)
-
- if "base_url" not in values:
- values["base_url"] = SOLAR_SERVICE_URL_BASE
-
- if "base_url" in values and not values["base_url"].startswith(SOLAR_SERVICE):
- raise ValueError("base_url must match with: " + SOLAR_SERVICE)
-
- values["_client"] = _SolarClient(
- api_key=values["solar_api_key"], base_url=values["base_url"]
- )
- return values
-
- @property
- def _llm_type(self) -> str:
- return "solar"
-
-
-class Solar(SolarCommon, LLM):
- """Solar large language models.
-
- To use, you should have the environment variable
- ``SOLAR_API_KEY`` set with your API key.
- Referenced from https://console.upstage.ai/services/solar
- """
-
- model_config = ConfigDict(
- populate_by_name=True,
- )
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- request = self._invocation_params
- request["messages"] = [{"role": "user", "content": prompt}]
- request.update(kwargs)
- text = self._client.completion(request)
- if stop is not None:
- # This is required since the stop tokens
- # are not enforced by the model parameters
- text = enforce_stop_tokens(text, stop)
-
- return text
diff --git a/libs/community/langchain_community/llms/sparkllm.py b/libs/community/langchain_community/llms/sparkllm.py
deleted file mode 100644
index 8f0ead4d27..0000000000
--- a/libs/community/langchain_community/llms/sparkllm.py
+++ /dev/null
@@ -1,469 +0,0 @@
-from __future__ import annotations
-
-import base64
-import hashlib
-import hmac
-import json
-import logging
-import queue
-import threading
-from datetime import datetime
-from queue import Queue
-from time import mktime
-from typing import Any, Dict, Generator, Iterator, List, Optional
-from urllib.parse import urlencode, urlparse, urlunparse
-from wsgiref.handlers import format_date_time
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import Field
-
-logger = logging.getLogger(__name__)
-
-
-class SparkLLM(LLM):
- """iFlyTek Spark completion model integration.
-
- Setup:
- To use, you should set environment variables ``IFLYTEK_SPARK_APP_ID``,
- ``IFLYTEK_SPARK_API_KEY`` and ``IFLYTEK_SPARK_API_SECRET``.
-
- .. code-block:: bash
-
- export IFLYTEK_SPARK_APP_ID="your-app-id"
- export IFLYTEK_SPARK_API_KEY="your-api-key"
- export IFLYTEK_SPARK_API_SECRET="your-api-secret"
-
- Key init args — completion params:
- model: Optional[str]
- Name of IFLYTEK SPARK model to use.
- temperature: Optional[float]
- Sampling temperature.
- top_k: Optional[float]
- What search sampling control to use.
- streaming: Optional[bool]
- Whether to stream the results or not.
-
- Key init args — client params:
- app_id: Optional[str]
- IFLYTEK SPARK API KEY. Automatically inferred from env var `IFLYTEK_SPARK_APP_ID` if not provided.
- api_key: Optional[str]
- IFLYTEK SPARK API KEY. If not passed in will be read from env var IFLYTEK_SPARK_API_KEY.
- api_secret: Optional[str]
- IFLYTEK SPARK API SECRET. If not passed in will be read from env var IFLYTEK_SPARK_API_SECRET.
- api_url: Optional[str]
- Base URL for API requests.
- timeout: Optional[int]
- Timeout for requests.
-
- See full list of supported init args and their descriptions in the params section.
-
- Instantiate:
- .. code-block:: python
-
- from langchain_community.llms import SparkLLM
-
- llm = SparkLLM(
- app_id="your-app-id",
- api_key="your-api_key",
- api_secret="your-api-secret",
- # model='Spark4.0 Ultra',
- # temperature=...,
- # other params...
- )
-
- Invoke:
- .. code-block:: python
-
- input_text = "用50个字左右阐述,生命的意义在于"
- llm.invoke(input_text)
-
- .. code-block:: python
-
- '生命的意义在于实现自我价值,追求内心的平静与快乐,同时为他人和社会带来正面影响。'
-
- Stream:
- .. code-block:: python
-
- for chunk in llm.stream(input_text):
- print(chunk)
-
- .. code-block:: python
-
- 生命 | 的意义在于 | 不断探索和 | 实现个人潜能,通过 | 学习 | 、成长和对社会 | 的贡献,追求内心的满足和幸福。
-
- Async:
- .. code-block:: python
-
- await llm.ainvoke(input_text)
-
- # stream:
- # async for chunk in llm.astream(input_text):
- # print(chunk)
-
- # batch:
- # await llm.abatch([input_text])
-
- .. code-block:: python
-
- '生命的意义在于实现自我价值,追求内心的平静与快乐,同时为他人和社会带来正面影响。'
-
- """ # noqa: E501
-
- client: Any = None #: :meta private:
- spark_app_id: Optional[str] = Field(default=None, alias="app_id")
- """Automatically inferred from env var `IFLYTEK_SPARK_APP_ID`
- if not provided."""
- spark_api_key: Optional[str] = Field(default=None, alias="api_key")
- """IFLYTEK SPARK API KEY. If not passed in will be read from
- env var IFLYTEK_SPARK_API_KEY."""
- spark_api_secret: Optional[str] = Field(default=None, alias="api_secret")
- """IFLYTEK SPARK API SECRET. If not passed in will be read from
- env var IFLYTEK_SPARK_API_SECRET."""
- spark_api_url: Optional[str] = Field(default=None, alias="api_url")
- """Base URL path for API requests, leave blank if not using a proxy or service
- emulator."""
- spark_llm_domain: Optional[str] = Field(default=None, alias="model")
- """Model name to use."""
- spark_user_id: str = "lc_user"
- streaming: bool = False
- """Whether to stream the results or not."""
- request_timeout: int = Field(default=30, alias="timeout")
- """request timeout for chat http requests"""
- temperature: float = 0.5
- """What sampling temperature to use."""
- top_k: int = 4
- """What search sampling control to use."""
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for API call not explicitly specified."""
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- values["spark_app_id"] = get_from_dict_or_env(
- values,
- ["spark_app_id", "app_id"],
- "IFLYTEK_SPARK_APP_ID",
- )
- values["spark_api_key"] = get_from_dict_or_env(
- values,
- ["spark_api_key", "api_key"],
- "IFLYTEK_SPARK_API_KEY",
- )
- values["spark_api_secret"] = get_from_dict_or_env(
- values,
- ["spark_api_secret", "api_secret"],
- "IFLYTEK_SPARK_API_SECRET",
- )
- values["spark_api_url"] = get_from_dict_or_env(
- values,
- ["spark_api_url", "api_url"],
- "IFLYTEK_SPARK_API_URL",
- "wss://spark-api.xf-yun.com/v3.5/chat",
- )
- values["spark_llm_domain"] = get_from_dict_or_env(
- values,
- ["spark_llm_domain", "model"],
- "IFLYTEK_SPARK_LLM_DOMAIN",
- "generalv3.5",
- )
- # put extra params into model_kwargs
- values["model_kwargs"]["temperature"] = values["temperature"] or cls.temperature
- values["model_kwargs"]["top_k"] = values["top_k"] or cls.top_k
-
- values["client"] = _SparkLLMClient(
- app_id=values["spark_app_id"],
- api_key=values["spark_api_key"],
- api_secret=values["spark_api_secret"],
- api_url=values["spark_api_url"],
- spark_domain=values["spark_llm_domain"],
- model_kwargs=values["model_kwargs"],
- )
- return values
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "spark-llm-chat"
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling SparkLLM API."""
- normal_params = {
- "spark_llm_domain": self.spark_llm_domain,
- "stream": self.streaming,
- "request_timeout": self.request_timeout,
- "top_k": self.top_k,
- "temperature": self.temperature,
- }
-
- return {**normal_params, **self.model_kwargs}
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to an sparkllm for each generation with a prompt.
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
- Returns:
- The string generated by the llm.
-
- Example:
- .. code-block:: python
- response = client("Tell me a joke.")
- """
- if self.streaming:
- completion = ""
- for chunk in self._stream(prompt, stop, run_manager, **kwargs):
- completion += chunk.text
- return completion
- completion = ""
- self.client.arun(
- [{"role": "user", "content": prompt}],
- self.spark_user_id,
- self.model_kwargs,
- self.streaming,
- )
- for content in self.client.subscribe(timeout=self.request_timeout):
- if "data" not in content:
- continue
- completion = content["data"]["content"]
-
- return completion
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- self.client.run(
- [{"role": "user", "content": prompt}],
- self.spark_user_id,
- self.model_kwargs,
- True,
- )
- for content in self.client.subscribe(timeout=self.request_timeout):
- if "data" not in content:
- continue
- delta = content["data"]
- if run_manager:
- run_manager.on_llm_new_token(delta)
- yield GenerationChunk(text=delta["content"])
-
-
-class _SparkLLMClient:
- """
- Use websocket-client to call the SparkLLM interface provided by Xfyun,
- which is the iFlyTek's open platform for AI capabilities
- """
-
- def __init__(
- self,
- app_id: str,
- api_key: str,
- api_secret: str,
- api_url: Optional[str] = None,
- spark_domain: Optional[str] = None,
- model_kwargs: Optional[dict] = None,
- ):
- try:
- import websocket
-
- self.websocket_client = websocket
- except ImportError:
- raise ImportError(
- "Could not import websocket client python package. "
- "Please install it with `pip install websocket-client`."
- )
-
- self.api_url = (
- "wss://spark-api.xf-yun.com/v3.5/chat" if not api_url else api_url
- )
- self.app_id = app_id
- self.model_kwargs = model_kwargs
- self.spark_domain = spark_domain or "generalv3.5"
- self.queue: Queue[Dict] = Queue()
- self.blocking_message = {"content": "", "role": "assistant"}
- self.api_key = api_key
- self.api_secret = api_secret
-
- @staticmethod
- def _create_url(api_url: str, api_key: str, api_secret: str) -> str:
- """
- Generate a request url with an api key and an api secret.
- """
- # generate timestamp by RFC1123
- date = format_date_time(mktime(datetime.now().timetuple()))
-
- # urlparse
- parsed_url = urlparse(api_url)
- host = parsed_url.netloc
- path = parsed_url.path
-
- signature_origin = f"host: {host}\ndate: {date}\nGET {path} HTTP/1.1"
-
- # encrypt using hmac-sha256
- signature_sha = hmac.new(
- api_secret.encode("utf-8"),
- signature_origin.encode("utf-8"),
- digestmod=hashlib.sha256,
- ).digest()
-
- signature_sha_base64 = base64.b64encode(signature_sha).decode(encoding="utf-8")
-
- authorization_origin = f'api_key="{api_key}", algorithm="hmac-sha256", \
- headers="host date request-line", signature="{signature_sha_base64}"'
- authorization = base64.b64encode(authorization_origin.encode("utf-8")).decode(
- encoding="utf-8"
- )
-
- # generate url
- params_dict = {"authorization": authorization, "date": date, "host": host}
- encoded_params = urlencode(params_dict)
- url = urlunparse(
- (
- parsed_url.scheme,
- parsed_url.netloc,
- parsed_url.path,
- parsed_url.params,
- encoded_params,
- parsed_url.fragment,
- )
- )
- return url
-
- def run(
- self,
- messages: List[Dict],
- user_id: str,
- model_kwargs: Optional[dict] = None,
- streaming: bool = False,
- ) -> None:
- self.websocket_client.enableTrace(False)
- ws = self.websocket_client.WebSocketApp(
- _SparkLLMClient._create_url(
- self.api_url,
- self.api_key,
- self.api_secret,
- ),
- on_message=self.on_message,
- on_error=self.on_error,
- on_close=self.on_close,
- on_open=self.on_open,
- )
- ws.messages = messages # type: ignore[attr-defined]
- ws.user_id = user_id # type: ignore[attr-defined]
- ws.model_kwargs = self.model_kwargs if model_kwargs is None else model_kwargs # type: ignore[attr-defined]
- ws.streaming = streaming # type: ignore[attr-defined]
- ws.run_forever()
-
- def arun(
- self,
- messages: List[Dict],
- user_id: str,
- model_kwargs: Optional[dict] = None,
- streaming: bool = False,
- ) -> threading.Thread:
- ws_thread = threading.Thread(
- target=self.run,
- args=(
- messages,
- user_id,
- model_kwargs,
- streaming,
- ),
- )
- ws_thread.start()
- return ws_thread
-
- def on_error(self, ws: Any, error: Optional[Any]) -> None:
- self.queue.put({"error": error})
- ws.close()
-
- def on_close(self, ws: Any, close_status_code: int, close_reason: str) -> None:
- logger.debug(
- {
- "log": {
- "close_status_code": close_status_code,
- "close_reason": close_reason,
- }
- }
- )
- self.queue.put({"done": True})
-
- def on_open(self, ws: Any) -> None:
- self.blocking_message = {"content": "", "role": "assistant"}
- data = json.dumps(
- self.gen_params(
- messages=ws.messages, user_id=ws.user_id, model_kwargs=ws.model_kwargs
- )
- )
- ws.send(data)
-
- def on_message(self, ws: Any, message: str) -> None:
- data = json.loads(message)
- code = data["header"]["code"]
- if code != 0:
- self.queue.put(
- {"error": f"Code: {code}, Error: {data['header']['message']}"}
- )
- ws.close()
- else:
- choices = data["payload"]["choices"]
- status = choices["status"]
- content = choices["text"][0]["content"]
- if ws.streaming:
- self.queue.put({"data": choices["text"][0]})
- else:
- self.blocking_message["content"] += content
- if status == 2:
- if not ws.streaming:
- self.queue.put({"data": self.blocking_message})
- usage_data = (
- data.get("payload", {}).get("usage", {}).get("text", {})
- if data
- else {}
- )
- self.queue.put({"usage": usage_data})
- ws.close()
-
- def gen_params(
- self, messages: list, user_id: str, model_kwargs: Optional[dict] = None
- ) -> dict:
- data: Dict = {
- "header": {"app_id": self.app_id, "uid": user_id},
- "parameter": {"chat": {"domain": self.spark_domain}},
- "payload": {"message": {"text": messages}},
- }
-
- if model_kwargs:
- data["parameter"]["chat"].update(model_kwargs)
- logger.debug(f"Spark Request Parameters: {data}")
- return data
-
- def subscribe(self, timeout: Optional[int] = 30) -> Generator[Dict, None, None]:
- while True:
- try:
- content = self.queue.get(timeout=timeout)
- except queue.Empty as _:
- raise TimeoutError(
- f"SparkLLMClient wait LLM api response timeout {timeout} seconds"
- )
- if "error" in content:
- raise ConnectionError(content["error"])
- if "usage" in content:
- yield content
- continue
- if "done" in content:
- break
- if "data" not in content:
- break
- yield content
diff --git a/libs/community/langchain_community/llms/stochasticai.py b/libs/community/langchain_community/llms/stochasticai.py
deleted file mode 100644
index 0e999bcb3a..0000000000
--- a/libs/community/langchain_community/llms/stochasticai.py
+++ /dev/null
@@ -1,137 +0,0 @@
-import logging
-import time
-from typing import Any, Dict, List, Mapping, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import ConfigDict, Field, SecretStr, model_validator
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-class StochasticAI(LLM):
- """StochasticAI large language models.
-
- To use, you should have the environment variable ``STOCHASTICAI_API_KEY``
- set with your API key.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import StochasticAI
- stochasticai = StochasticAI(api_url="")
- """
-
- api_url: str = ""
- """Model name to use."""
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `create` call not
- explicitly specified."""
-
- stochasticai_api_key: Optional[SecretStr] = None
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def build_extra(cls, values: Dict[str, Any]) -> Any:
- """Build extra kwargs from additional params that were passed in."""
- all_required_field_names = set(list(cls.model_fields.keys()))
-
- extra = values.get("model_kwargs", {})
- for field_name in list(values):
- if field_name not in all_required_field_names:
- if field_name in extra:
- raise ValueError(f"Found {field_name} supplied twice.")
- logger.warning(
- f"""{field_name} was transferred to model_kwargs.
- Please confirm that {field_name} is what you intended."""
- )
- extra[field_name] = values.pop(field_name)
- values["model_kwargs"] = extra
- return values
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key exists in environment."""
- stochasticai_api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "stochasticai_api_key", "STOCHASTICAI_API_KEY")
- )
- values["stochasticai_api_key"] = stochasticai_api_key
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- **{"endpoint_url": self.api_url},
- **{"model_kwargs": self.model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "stochasticai"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to StochasticAI's complete endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = StochasticAI("Tell me a joke.")
- """
- params = self.model_kwargs or {}
- params = {**params, **kwargs}
- response_post = requests.post(
- url=self.api_url,
- json={"prompt": prompt, "params": params},
- headers={
- "apiKey": f"{self.stochasticai_api_key.get_secret_value()}", # type: ignore[union-attr]
- "Accept": "application/json",
- "Content-Type": "application/json",
- },
- )
- response_post.raise_for_status()
- response_post_json = response_post.json()
- completed = False
- while not completed:
- response_get = requests.get(
- url=response_post_json["data"]["responseUrl"],
- headers={
- "apiKey": f"{self.stochasticai_api_key.get_secret_value()}", # type: ignore[union-attr]
- "Accept": "application/json",
- "Content-Type": "application/json",
- },
- )
- response_get.raise_for_status()
- response_get_json = response_get.json()["data"]
- text = response_get_json.get("completion")
- completed = text is not None
- time.sleep(0.5)
- text = text[0]
- if stop is not None:
- # I believe this is required since the stop tokens
- # are not enforced by the model parameters
- text = enforce_stop_tokens(text, stop)
- return text
diff --git a/libs/community/langchain_community/llms/symblai_nebula.py b/libs/community/langchain_community/llms/symblai_nebula.py
deleted file mode 100644
index 63bb29a9a5..0000000000
--- a/libs/community/langchain_community/llms/symblai_nebula.py
+++ /dev/null
@@ -1,231 +0,0 @@
-import json
-import logging
-from typing import Any, Callable, Dict, List, Mapping, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import ConfigDict, SecretStr
-from requests import ConnectTimeout, ReadTimeout, RequestException
-from tenacity import (
- before_sleep_log,
- retry,
- retry_if_exception_type,
- stop_after_attempt,
- wait_exponential,
-)
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-DEFAULT_NEBULA_SERVICE_URL = "https://api-nebula.symbl.ai"
-DEFAULT_NEBULA_SERVICE_PATH = "/v1/model/generate"
-
-logger = logging.getLogger(__name__)
-
-
-class Nebula(LLM):
- """Nebula Service models.
-
- To use, you should have the environment variable ``NEBULA_SERVICE_URL``,
- ``NEBULA_SERVICE_PATH`` and ``NEBULA_API_KEY`` set with your Nebula
- Service, or pass it as a named parameter to the constructor.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import Nebula
-
- nebula = Nebula(
- nebula_service_url="NEBULA_SERVICE_URL",
- nebula_service_path="NEBULA_SERVICE_PATH",
- nebula_api_key="NEBULA_API_KEY",
- )
- """
-
- """Key/value arguments to pass to the model. Reserved for future use"""
- model_kwargs: Optional[dict] = None
-
- """Optional"""
-
- nebula_service_url: Optional[str] = None
- nebula_service_path: Optional[str] = None
- nebula_api_key: Optional[SecretStr] = None
- model: Optional[str] = None
- max_new_tokens: Optional[int] = 128
- temperature: Optional[float] = 0.6
- top_p: Optional[float] = 0.95
- repetition_penalty: Optional[float] = 1.0
- top_k: Optional[int] = 1
- stop_sequences: Optional[List[str]] = None
- max_retries: Optional[int] = 10
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- nebula_service_url = get_from_dict_or_env(
- values,
- "nebula_service_url",
- "NEBULA_SERVICE_URL",
- DEFAULT_NEBULA_SERVICE_URL,
- )
- nebula_service_path = get_from_dict_or_env(
- values,
- "nebula_service_path",
- "NEBULA_SERVICE_PATH",
- DEFAULT_NEBULA_SERVICE_PATH,
- )
- nebula_api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "nebula_api_key", "NEBULA_API_KEY", None)
- )
-
- if nebula_service_url.endswith("/"):
- nebula_service_url = nebula_service_url[:-1]
- if not nebula_service_path.startswith("/"):
- nebula_service_path = "/" + nebula_service_path
-
- values["nebula_service_url"] = nebula_service_url
- values["nebula_service_path"] = nebula_service_path
- values["nebula_api_key"] = nebula_api_key
-
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling Cohere API."""
- return {
- "max_new_tokens": self.max_new_tokens,
- "temperature": self.temperature,
- "top_k": self.top_k,
- "top_p": self.top_p,
- "repetition_penalty": self.repetition_penalty,
- }
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- _model_kwargs = self.model_kwargs or {}
- return {
- "nebula_service_url": self.nebula_service_url,
- "nebula_service_path": self.nebula_service_path,
- **{"model_kwargs": _model_kwargs},
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "nebula"
-
- def _invocation_params(
- self, stop_sequences: Optional[List[str]], **kwargs: Any
- ) -> dict:
- params = self._default_params
- if self.stop_sequences is not None and stop_sequences is not None:
- raise ValueError("`stop` found in both the input and default params.")
- elif self.stop_sequences is not None:
- params["stop_sequences"] = self.stop_sequences
- else:
- params["stop_sequences"] = stop_sequences
- return {**params, **kwargs}
-
- @staticmethod
- def _process_response(response: Any, stop: Optional[List[str]]) -> str:
- text = response["output"]["text"]
- if stop:
- text = enforce_stop_tokens(text, stop)
- return text
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Nebula Service endpoint.
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
- Returns:
- The string generated by the model.
- Example:
- .. code-block:: python
- response = nebula("Tell me a joke.")
- """
- params = self._invocation_params(stop, **kwargs)
- prompt = prompt.strip()
-
- response = completion_with_retry(
- self,
- prompt=prompt,
- params=params,
- url=f"{self.nebula_service_url}{self.nebula_service_path}",
- )
- _stop = params.get("stop_sequences")
- return self._process_response(response, _stop)
-
-
-def make_request(
- self: Nebula,
- prompt: str,
- url: str = f"{DEFAULT_NEBULA_SERVICE_URL}{DEFAULT_NEBULA_SERVICE_PATH}",
- params: Optional[Dict] = None,
-) -> Any:
- """Generate text from the model."""
- params = params or {}
- api_key = None
- if self.nebula_api_key is not None:
- api_key = self.nebula_api_key.get_secret_value()
- headers = {
- "Content-Type": "application/json",
- "ApiKey": f"{api_key}",
- }
-
- body = {"prompt": prompt}
-
- # add params to body
- for key, value in params.items():
- body[key] = value
-
- # make request
- response = requests.post(url, headers=headers, json=body)
-
- if response.status_code != 200:
- raise Exception(
- f"Request failed with status code {response.status_code}"
- f" and message {response.text}"
- )
-
- return json.loads(response.text)
-
-
-def _create_retry_decorator(llm: Nebula) -> Callable[[Any], Any]:
- min_seconds = 4
- max_seconds = 10
- # Wait 2^x * 1 second between each retry starting with
- # 4 seconds, then up to 10 seconds, then 10 seconds afterward
- max_retries = llm.max_retries if llm.max_retries is not None else 3
- return retry(
- reraise=True,
- stop=stop_after_attempt(max_retries),
- wait=wait_exponential(multiplier=1, min=min_seconds, max=max_seconds),
- retry=(
- retry_if_exception_type((RequestException, ConnectTimeout, ReadTimeout))
- ),
- before_sleep=before_sleep_log(logger, logging.WARNING),
- )
-
-
-def completion_with_retry(llm: Nebula, **kwargs: Any) -> Any:
- """Use tenacity to retry the completion call."""
- retry_decorator = _create_retry_decorator(llm)
-
- @retry_decorator
- def _completion_with_retry(**_kwargs: Any) -> Any:
- return make_request(llm, **_kwargs)
-
- return _completion_with_retry(**kwargs)
diff --git a/libs/community/langchain_community/llms/textgen.py b/libs/community/langchain_community/llms/textgen.py
deleted file mode 100644
index 33a4f74f97..0000000000
--- a/libs/community/langchain_community/llms/textgen.py
+++ /dev/null
@@ -1,415 +0,0 @@
-import json
-import logging
-from typing import Any, AsyncIterator, Dict, Iterator, List, Optional
-
-import requests
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from pydantic import Field
-
-logger = logging.getLogger(__name__)
-
-
-class TextGen(LLM):
- """Text generation models from WebUI.
-
- To use, you should have the text-generation-webui installed, a model loaded,
- and --api added as a command-line option.
-
- Suggested installation, use one-click installer for your OS:
- https://github.com/oobabooga/text-generation-webui#one-click-installers
-
- Parameters below taken from text-generation-webui api example:
- https://github.com/oobabooga/text-generation-webui/blob/main/api-examples/api-example.py
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import TextGen
- llm = TextGen(model_url="http://localhost:8500")
- """
-
- model_url: str
- """The full URL to the textgen webui including http[s]://host:port """
-
- preset: Optional[str] = None
- """The preset to use in the textgen webui """
-
- max_new_tokens: Optional[int] = 250
- """The maximum number of tokens to generate."""
-
- do_sample: bool = Field(True, alias="do_sample")
- """Do sample"""
-
- temperature: Optional[float] = 1.3
- """Primary factor to control randomness of outputs. 0 = deterministic
- (only the most likely token is used). Higher value = more randomness."""
-
- top_p: Optional[float] = 0.1
- """If not set to 1, select tokens with probabilities adding up to less than this
- number. Higher value = higher range of possible random results."""
-
- typical_p: Optional[float] = 1
- """If not set to 1, select only tokens that are at least this much more likely to
- appear than random tokens, given the prior text."""
-
- epsilon_cutoff: Optional[float] = 0 # In units of 1e-4
- """Epsilon cutoff"""
-
- eta_cutoff: Optional[float] = 0 # In units of 1e-4
- """ETA cutoff"""
-
- repetition_penalty: Optional[float] = 1.18
- """Exponential penalty factor for repeating prior tokens. 1 means no penalty,
- higher value = less repetition, lower value = more repetition."""
-
- top_k: Optional[float] = 40
- """Similar to top_p, but select instead only the top_k most likely tokens.
- Higher value = higher range of possible random results."""
-
- min_length: Optional[int] = 0
- """Minimum generation length in tokens."""
-
- no_repeat_ngram_size: Optional[int] = 0
- """If not set to 0, specifies the length of token sets that are completely blocked
- from repeating at all. Higher values = blocks larger phrases,
- lower values = blocks words or letters from repeating.
- Only 0 or high values are a good idea in most cases."""
-
- num_beams: Optional[int] = 1
- """Number of beams"""
-
- penalty_alpha: Optional[float] = 0
- """Penalty Alpha"""
-
- length_penalty: Optional[float] = 1
- """Length Penalty"""
-
- early_stopping: bool = Field(False, alias="early_stopping")
- """Early stopping"""
-
- seed: int = Field(-1, alias="seed")
- """Seed (-1 for random)"""
-
- add_bos_token: bool = Field(True, alias="add_bos_token")
- """Add the bos_token to the beginning of prompts.
- Disabling this can make the replies more creative."""
-
- truncation_length: Optional[int] = 2048
- """Truncate the prompt up to this length. The leftmost tokens are removed if
- the prompt exceeds this length. Most models require this to be at most 2048."""
-
- ban_eos_token: bool = Field(False, alias="ban_eos_token")
- """Ban the eos_token. Forces the model to never end the generation prematurely."""
-
- skip_special_tokens: bool = Field(True, alias="skip_special_tokens")
- """Skip special tokens. Some specific models need this unset."""
-
- stopping_strings: Optional[List[str]] = []
- """A list of strings to stop generation when encountered."""
-
- streaming: bool = False
- """Whether to stream the results, token by token."""
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling textgen."""
- return {
- "max_new_tokens": self.max_new_tokens,
- "do_sample": self.do_sample,
- "temperature": self.temperature,
- "top_p": self.top_p,
- "typical_p": self.typical_p,
- "epsilon_cutoff": self.epsilon_cutoff,
- "eta_cutoff": self.eta_cutoff,
- "repetition_penalty": self.repetition_penalty,
- "top_k": self.top_k,
- "min_length": self.min_length,
- "no_repeat_ngram_size": self.no_repeat_ngram_size,
- "num_beams": self.num_beams,
- "penalty_alpha": self.penalty_alpha,
- "length_penalty": self.length_penalty,
- "early_stopping": self.early_stopping,
- "seed": self.seed,
- "add_bos_token": self.add_bos_token,
- "truncation_length": self.truncation_length,
- "ban_eos_token": self.ban_eos_token,
- "skip_special_tokens": self.skip_special_tokens,
- "stopping_strings": self.stopping_strings,
- }
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {**{"model_url": self.model_url}, **self._default_params}
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "textgen"
-
- def _get_parameters(self, stop: Optional[List[str]] = None) -> Dict[str, Any]:
- """
- Performs sanity check, preparing parameters in format needed by textgen.
-
- Args:
- stop (Optional[List[str]]): List of stop sequences for textgen.
-
- Returns:
- Dictionary containing the combined parameters.
- """
-
- # Raise error if stop sequences are in both input and default params
- # if self.stop and stop is not None:
- if self.stopping_strings and stop is not None:
- raise ValueError("`stop` found in both the input and default params.")
-
- if self.preset is None:
- params = self._default_params
- else:
- params = {"preset": self.preset}
-
- # then sets it as configured, or default to an empty list:
- params["stopping_strings"] = self.stopping_strings or stop or []
-
- return params
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call the textgen web API and return the output.
-
- Args:
- prompt: The prompt to use for generation.
- stop: A list of strings to stop generation when encountered.
-
- Returns:
- The generated text.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import TextGen
- llm = TextGen(model_url="http://localhost:5000")
- llm.invoke("Write a story about llamas.")
- """
- if self.streaming:
- combined_text_output = ""
- for chunk in self._stream(
- prompt=prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- combined_text_output += chunk.text
- result = combined_text_output
-
- else:
- url = f"{self.model_url}/api/v1/generate"
- params = self._get_parameters(stop)
- request = params.copy()
- request["prompt"] = prompt
- response = requests.post(url, json=request)
-
- if response.status_code == 200:
- result = response.json()["results"][0]["text"]
- else:
- print(f"ERROR: response: {response}") # noqa: T201
- result = ""
-
- return result
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call the textgen web API and return the output.
-
- Args:
- prompt: The prompt to use for generation.
- stop: A list of strings to stop generation when encountered.
-
- Returns:
- The generated text.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import TextGen
- llm = TextGen(model_url="http://localhost:5000")
- llm.invoke("Write a story about llamas.")
- """
- if self.streaming:
- combined_text_output = ""
- async for chunk in self._astream(
- prompt=prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- combined_text_output += chunk.text
- result = combined_text_output
-
- else:
- url = f"{self.model_url}/api/v1/generate"
- params = self._get_parameters(stop)
- request = params.copy()
- request["prompt"] = prompt
- response = requests.post(url, json=request)
-
- if response.status_code == 200:
- result = response.json()["results"][0]["text"]
- else:
- print(f"ERROR: response: {response}") # noqa: T201
- result = ""
-
- return result
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- """Yields results objects as they are generated in real time.
-
- It also calls the callback manager's on_llm_new_token event with
- similar parameters to the OpenAI LLM class method of the same name.
-
- Args:
- prompt: The prompts to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- A generator representing the stream of tokens being generated.
-
- Yields:
- A dictionary like objects containing a string token and metadata.
- See text-generation-webui docs and below for more.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import TextGen
- llm = TextGen(
- model_url = "ws://localhost:5005"
- streaming=True
- )
- for chunk in llm.stream("Ask 'Hi, how are you?' like a pirate:'",
- stop=["'","\n"]):
- print(chunk, end='', flush=True) # noqa: T201
-
- """
- try:
- import websocket
- except ImportError:
- raise ImportError(
- "The `websocket-client` package is required for streaming."
- )
-
- params = {**self._get_parameters(stop), **kwargs}
-
- url = f"{self.model_url}/api/v1/stream"
-
- request = params.copy()
- request["prompt"] = prompt
-
- websocket_client = websocket.WebSocket()
-
- websocket_client.connect(url)
-
- websocket_client.send(json.dumps(request))
-
- while True:
- result = websocket_client.recv()
- result = json.loads(result)
-
- if result["event"] == "text_stream": # type: ignore[call-overload, index]
- chunk = GenerationChunk(
- text=result["text"], # type: ignore[call-overload, index]
- generation_info=None,
- )
- if run_manager:
- run_manager.on_llm_new_token(token=chunk.text)
- yield chunk
- elif result["event"] == "stream_end": # type: ignore[call-overload, index]
- websocket_client.close()
- return
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- """Yields results objects as they are generated in real time.
-
- It also calls the callback manager's on_llm_new_token event with
- similar parameters to the OpenAI LLM class method of the same name.
-
- Args:
- prompt: The prompts to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- A generator representing the stream of tokens being generated.
-
- Yields:
- A dictionary like objects containing a string token and metadata.
- See text-generation-webui docs and below for more.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import TextGen
- llm = TextGen(
- model_url = "ws://localhost:5005"
- streaming=True
- )
- for chunk in llm.stream("Ask 'Hi, how are you?' like a pirate:'",
- stop=["'","\n"]):
- print(chunk, end='', flush=True) # noqa: T201
-
- """
- try:
- import websocket
- except ImportError:
- raise ImportError(
- "The `websocket-client` package is required for streaming."
- )
-
- params = {**self._get_parameters(stop), **kwargs}
-
- url = f"{self.model_url}/api/v1/stream"
-
- request = params.copy()
- request["prompt"] = prompt
-
- websocket_client = websocket.WebSocket()
-
- websocket_client.connect(url)
-
- websocket_client.send(json.dumps(request))
-
- while True:
- result = websocket_client.recv()
- result = json.loads(result)
-
- if result["event"] == "text_stream": # type: ignore[call-overload, index]
- chunk = GenerationChunk(
- text=result["text"], # type: ignore[call-overload, index]
- generation_info=None,
- )
- if run_manager:
- await run_manager.on_llm_new_token(token=chunk.text)
- yield chunk
- elif result["event"] == "stream_end": # type: ignore[call-overload, index]
- websocket_client.close()
- return
diff --git a/libs/community/langchain_community/llms/titan_takeoff.py b/libs/community/langchain_community/llms/titan_takeoff.py
deleted file mode 100644
index 7f1d765a0d..0000000000
--- a/libs/community/langchain_community/llms/titan_takeoff.py
+++ /dev/null
@@ -1,264 +0,0 @@
-from enum import Enum
-from typing import Any, Iterator, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from pydantic import BaseModel, ConfigDict
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-
-class Device(str, Enum):
- """The device to use for inference, cuda or cpu"""
-
- cuda = "cuda"
- cpu = "cpu"
-
-
-class ReaderConfig(BaseModel):
- """Configuration for the reader to be deployed in Titan Takeoff API."""
-
- model_config = ConfigDict(
- protected_namespaces=(),
- )
-
- model_name: str
- """The name of the model to use"""
-
- device: Device = Device.cuda
- """The device to use for inference, cuda or cpu"""
-
- consumer_group: str = "primary"
- """The consumer group to place the reader into"""
-
- tensor_parallel: Optional[int] = None
- """The number of gpus you would like your model to be split across"""
-
- max_seq_length: int = 512
- """The maximum sequence length to use for inference, defaults to 512"""
-
- max_batch_size: int = 4
- """The max batch size for continuous batching of requests"""
-
-
-class TitanTakeoff(LLM):
- """Titan Takeoff API LLMs.
-
- Titan Takeoff is a wrapper to interface with Takeoff Inference API for
- generative text to text language models.
-
- You can use this wrapper to send requests to a generative language model
- and to deploy readers with Takeoff.
-
- Examples:
- This is an example how to deploy a generative language model and send
- requests.
-
- .. code-block:: python
- # Import the TitanTakeoff class from community package
- import time
- from langchain_community.llms import TitanTakeoff
-
- # Specify the embedding reader you'd like to deploy
- reader_1 = {
- "model_name": "TheBloke/Llama-2-7b-Chat-AWQ",
- "device": "cuda",
- "tensor_parallel": 1,
- "consumer_group": "llama"
- }
-
- # For every reader you pass into models arg Takeoff will spin
- # up a reader according to the specs you provide. If you don't
- # specify the arg no models are spun up and it assumes you have
- # already done this separately.
- llm = TitanTakeoff(models=[reader_1])
-
- # Wait for the reader to be deployed, time needed depends on the
- # model size and your internet speed
- time.sleep(60)
-
- # Returns the query, ie a List[float], sent to `llama` consumer group
- # where we just spun up the Llama 7B model
- print(embed.invoke(
- "Where can I see football?", consumer_group="llama"
- ))
-
- # You can also send generation parameters to the model, any of the
- # following can be passed in as kwargs:
- # https://docs.titanml.co/docs/next/apis/Takeoff%20inference_REST_API/generate#request
- # for instance:
- print(embed.invoke(
- "Where can I see football?", consumer_group="llama", max_new_tokens=100
- ))
- """
-
- base_url: str = "http://localhost"
- """The base URL of the Titan Takeoff (Pro) server. Default = "http://localhost"."""
-
- port: int = 3000
- """The port of the Titan Takeoff (Pro) server. Default = 3000."""
-
- mgmt_port: int = 3001
- """The management port of the Titan Takeoff (Pro) server. Default = 3001."""
-
- streaming: bool = False
- """Whether to stream the output. Default = False."""
-
- client: Any = None
- """Takeoff Client Python SDK used to interact with Takeoff API"""
-
- def __init__(
- self,
- base_url: str = "http://localhost",
- port: int = 3000,
- mgmt_port: int = 3001,
- streaming: bool = False,
- models: List[ReaderConfig] = [],
- ):
- """Initialize the Titan Takeoff language wrapper.
-
- Args:
- base_url (str, optional): The base URL where the Takeoff
- Inference Server is listening. Defaults to `http://localhost`.
- port (int, optional): What port is Takeoff Inference API
- listening on. Defaults to 3000.
- mgmt_port (int, optional): What port is Takeoff Management API
- listening on. Defaults to 3001.
- streaming (bool, optional): Whether you want to by default use the
- generate_stream endpoint over generate to stream responses.
- Defaults to False. In reality, this is not significantly different
- as the streamed response is buffered and returned similar to the
- non-streamed response, but the run manager is applied per token
- generated.
- models (List[ReaderConfig], optional): Any readers you'd like to
- spin up on. Defaults to [].
-
- Raises:
- ImportError: If you haven't installed takeoff-client, you will
- get an ImportError. To remedy run `pip install 'takeoff-client==0.4.0'`
- """
- super().__init__( # type: ignore[call-arg]
- base_url=base_url, port=port, mgmt_port=mgmt_port, streaming=streaming
- )
- try:
- from takeoff_client import TakeoffClient
- except ImportError:
- raise ImportError(
- "takeoff-client is required for TitanTakeoff. "
- "Please install it with `pip install 'takeoff-client>=0.4.0'`."
- )
- self.client = TakeoffClient(
- self.base_url, port=self.port, mgmt_port=self.mgmt_port
- )
- for model in models:
- self.client.create_reader(model)
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "titan_takeoff"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Titan Takeoff (Pro) generate endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
- run_manager: Optional callback manager to use when streaming.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- model = TitanTakeoff()
-
- prompt = "What is the capital of the United Kingdom?"
-
- # Use of model(prompt), ie `__call__` was deprecated in LangChain 0.1.7,
- # use model.invoke(prompt) instead.
- response = model.invoke(prompt)
-
- """
- if self.streaming:
- text_output = ""
- for chunk in self._stream(
- prompt=prompt,
- stop=stop,
- run_manager=run_manager,
- ):
- text_output += chunk.text
- return text_output
-
- response = self.client.generate(prompt, **kwargs)
- text = response["text"]
-
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
- return text
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- """Call out to Titan Takeoff (Pro) stream endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
- run_manager: Optional callback manager to use when streaming.
-
- Yields:
- A dictionary like object containing a string token.
-
- Example:
- .. code-block:: python
-
- model = TitanTakeoff()
-
- prompt = "What is the capital of the United Kingdom?"
- response = model.stream(prompt)
-
- # OR
-
- model = TitanTakeoff(streaming=True)
-
- response = model.invoke(prompt)
-
- """
- response = self.client.generate_stream(prompt, **kwargs)
- buffer = ""
- for text in response:
- buffer += text.data
- if "data:" in buffer:
- # Remove the first instance of "data:" from the buffer.
- if buffer.startswith("data:"):
- buffer = ""
- if len(buffer.split("data:", 1)) == 2:
- content, _ = buffer.split("data:", 1)
- buffer = content.rstrip("\n")
- # Trim the buffer to only have content after the "data:" part.
- if buffer: # Ensure that there's content to process.
- chunk = GenerationChunk(text=buffer)
- buffer = "" # Reset buffer for the next set of data.
- if run_manager:
- run_manager.on_llm_new_token(token=chunk.text)
- yield chunk
-
- # Yield any remaining content in the buffer.
- if buffer:
- chunk = GenerationChunk(text=buffer.replace("", ""))
- if run_manager:
- run_manager.on_llm_new_token(token=chunk.text)
- yield chunk
diff --git a/libs/community/langchain_community/llms/together.py b/libs/community/langchain_community/llms/together.py
deleted file mode 100644
index e5e7b8d68b..0000000000
--- a/libs/community/langchain_community/llms/together.py
+++ /dev/null
@@ -1,211 +0,0 @@
-"""Wrapper around Together AI's Completion API."""
-
-import logging
-from typing import Any, Dict, List, Optional
-
-from aiohttp import ClientSession
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env
-from pydantic import ConfigDict, SecretStr, model_validator
-
-from langchain_community.utilities.requests import Requests
-
-logger = logging.getLogger(__name__)
-
-
-@deprecated(
- since="0.0.12", removal="1.0", alternative_import="langchain_together.Together"
-)
-class Together(LLM):
- """LLM models from `Together`.
-
- To use, you'll need an API key which you can find here:
- https://api.together.xyz/settings/api-keys. This can be passed in as init param
- ``together_api_key`` or set as environment variable ``TOGETHER_API_KEY``.
-
- Together AI API reference: https://docs.together.ai/reference/inference
- """
-
- base_url: str = "https://api.together.xyz/inference"
- """Base inference API URL."""
- together_api_key: SecretStr
- """Together AI API key. Get it here: https://api.together.xyz/settings/api-keys"""
- model: str
- """Model name. Available models listed here:
- https://docs.together.ai/docs/inference-models
- """
- temperature: Optional[float] = None
- """Model temperature."""
- top_p: Optional[float] = None
- """Used to dynamically adjust the number of choices for each predicted token based
- on the cumulative probabilities. A value of 1 will always yield the same
- output. A temperature less than 1 favors more correctness and is appropriate
- for question answering or summarization. A value greater than 1 introduces more
- randomness in the output.
- """
- top_k: Optional[int] = None
- """Used to limit the number of choices for the next predicted word or token. It
- specifies the maximum number of tokens to consider at each step, based on their
- probability of occurrence. This technique helps to speed up the generation
- process and can improve the quality of the generated text by focusing on the
- most likely options.
- """
- max_tokens: Optional[int] = None
- """The maximum number of tokens to generate."""
- repetition_penalty: Optional[float] = None
- """A number that controls the diversity of generated text by reducing the
- likelihood of repeated sequences. Higher values decrease repetition.
- """
- logprobs: Optional[int] = None
- """An integer that specifies how many top token log probabilities are included in
- the response for each token generation step.
- """
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key exists in environment."""
- values["together_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(values, "together_api_key", "TOGETHER_API_KEY")
- )
- return values
-
- @property
- def _llm_type(self) -> str:
- """Return type of model."""
- return "together"
-
- def _format_output(self, output: dict) -> str:
- return output["output"]["choices"][0]["text"]
-
- @staticmethod
- def get_user_agent() -> str:
- from langchain_community import __version__
-
- return f"langchain/{__version__}"
-
- @property
- def default_params(self) -> Dict[str, Any]:
- return {
- "model": self.model,
- "temperature": self.temperature,
- "top_p": self.top_p,
- "top_k": self.top_k,
- "max_tokens": self.max_tokens,
- "repetition_penalty": self.repetition_penalty,
- }
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to Together's text generation endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
-
- Returns:
- The string generated by the model..
- """
-
- headers = {
- "Authorization": f"Bearer {self.together_api_key.get_secret_value()}",
- "Content-Type": "application/json",
- }
- stop_to_use = stop[0] if stop and len(stop) == 1 else stop
- payload: Dict[str, Any] = {
- **self.default_params,
- "prompt": prompt,
- "stop": stop_to_use,
- **kwargs,
- }
-
- # filter None values to not pass them to the http payload
- payload = {k: v for k, v in payload.items() if v is not None}
- request = Requests(headers=headers)
- response = request.post(url=self.base_url, data=payload)
-
- if response.status_code >= 500:
- raise Exception(f"Together Server: Error {response.status_code}")
- elif response.status_code >= 400:
- raise ValueError(f"Together received an invalid payload: {response.text}")
- elif response.status_code != 200:
- raise Exception(
- f"Together returned an unexpected response with status "
- f"{response.status_code}: {response.text}"
- )
-
- data = response.json()
- if data.get("status") != "finished":
- err_msg = data.get("error", "Undefined Error")
- raise Exception(err_msg)
-
- output = self._format_output(data)
-
- return output
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call Together model to get predictions based on the prompt.
-
- Args:
- prompt: The prompt to pass into the model.
-
- Returns:
- The string generated by the model.
- """
- headers = {
- "Authorization": f"Bearer {self.together_api_key.get_secret_value()}",
- "Content-Type": "application/json",
- }
- stop_to_use = stop[0] if stop and len(stop) == 1 else stop
- payload: Dict[str, Any] = {
- **self.default_params,
- "prompt": prompt,
- "stop": stop_to_use,
- **kwargs,
- }
-
- # filter None values to not pass them to the http payload
- payload = {k: v for k, v in payload.items() if v is not None}
- async with ClientSession() as session:
- async with session.post(
- self.base_url, json=payload, headers=headers
- ) as response:
- if response.status >= 500:
- raise Exception(f"Together Server: Error {response.status}")
- elif response.status >= 400:
- raise ValueError(
- f"Together received an invalid payload: {response.text}"
- )
- elif response.status != 200:
- raise Exception(
- f"Together returned an unexpected response with status "
- f"{response.status}: {response.text}"
- )
-
- response_json = await response.json()
-
- if response_json.get("status") != "finished":
- err_msg = response_json.get("error", "Undefined Error")
- raise Exception(err_msg)
-
- output = self._format_output(response_json)
- return output
diff --git a/libs/community/langchain_community/llms/tongyi.py b/libs/community/langchain_community/llms/tongyi.py
deleted file mode 100644
index ade4d502a3..0000000000
--- a/libs/community/langchain_community/llms/tongyi.py
+++ /dev/null
@@ -1,462 +0,0 @@
-from __future__ import annotations
-
-import asyncio
-import functools
-import logging
-from typing import (
- Any,
- AsyncIterable,
- AsyncIterator,
- Callable,
- Dict,
- Iterable,
- Iterator,
- List,
- Mapping,
- Optional,
- Tuple,
- TypeVar,
-)
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import BaseLLM
-from langchain_core.outputs import Generation, GenerationChunk, LLMResult
-from langchain_core.utils import get_from_dict_or_env, pre_init
-from pydantic import Field
-from requests.exceptions import HTTPError
-from tenacity import (
- before_sleep_log,
- retry,
- retry_if_exception_type,
- stop_after_attempt,
- wait_exponential,
-)
-
-logger = logging.getLogger(__name__)
-T = TypeVar("T")
-
-
-def _create_retry_decorator(llm: Tongyi) -> Callable[[Any], Any]:
- min_seconds = 1
- max_seconds = 4
- # Wait 2^x * 1 second between each retry starting with
- # 4 seconds, then up to 10 seconds, then 10 seconds afterward
- return retry(
- reraise=True,
- stop=stop_after_attempt(llm.max_retries),
- wait=wait_exponential(multiplier=1, min=min_seconds, max=max_seconds),
- retry=(retry_if_exception_type(HTTPError)),
- before_sleep=before_sleep_log(logger, logging.WARNING),
- )
-
-
-def check_response(resp: Any) -> Any:
- """Check the response from the completion call."""
- if resp["status_code"] == 200:
- return resp
- elif resp["status_code"] in [400, 401]:
- raise ValueError(
- f"request_id: {resp['request_id']} \n "
- f"status_code: {resp['status_code']} \n "
- f"code: {resp['code']} \n message: {resp['message']}"
- )
- else:
- raise HTTPError(
- f"HTTP error occurred: status_code: {resp['status_code']} \n "
- f"code: {resp['code']} \n message: {resp['message']}",
- response=resp,
- )
-
-
-def generate_with_retry(llm: Tongyi, **kwargs: Any) -> Any:
- """Use tenacity to retry the completion call."""
- retry_decorator = _create_retry_decorator(llm)
-
- @retry_decorator
- def _generate_with_retry(**_kwargs: Any) -> Any:
- resp = llm.client.call(**_kwargs)
- return check_response(resp)
-
- return _generate_with_retry(**kwargs)
-
-
-def stream_generate_with_retry(llm: Tongyi, **kwargs: Any) -> Any:
- """Use tenacity to retry the completion call."""
- retry_decorator = _create_retry_decorator(llm)
-
- @retry_decorator
- def _stream_generate_with_retry(**_kwargs: Any) -> Any:
- responses = llm.client.call(**_kwargs)
- for resp in responses:
- yield check_response(resp)
-
- return _stream_generate_with_retry(**kwargs)
-
-
-async def astream_generate_with_retry(llm: Tongyi, **kwargs: Any) -> Any:
- """Async version of `stream_generate_with_retry`.
-
- Because the dashscope SDK doesn't provide an async API,
- we wrap `stream_generate_with_retry` with an async generator."""
-
- class _AioTongyiGenerator:
- def __init__(self, _llm: Tongyi, **_kwargs: Any):
- self.generator = stream_generate_with_retry(_llm, **_kwargs)
-
- def __aiter__(self) -> AsyncIterator[Any]:
- return self
-
- async def __anext__(self) -> Any:
- value = await asyncio.get_running_loop().run_in_executor(
- None, self._safe_next
- )
- if value is not None:
- return value
- else:
- raise StopAsyncIteration
-
- def _safe_next(self) -> Any:
- try:
- return next(self.generator)
- except StopIteration:
- return None
-
- async for chunk in _AioTongyiGenerator(llm, **kwargs):
- yield chunk
-
-
-def generate_with_last_element_mark(iterable: Iterable[T]) -> Iterator[Tuple[T, bool]]:
- """Generate elements from an iterable,
- and a boolean indicating if it is the last element."""
- iterator = iter(iterable)
- try:
- item = next(iterator)
- except StopIteration:
- return
- for next_item in iterator:
- yield item, False
- item = next_item
- yield item, True
-
-
-async def agenerate_with_last_element_mark(
- iterable: AsyncIterable[T],
-) -> AsyncIterator[Tuple[T, bool]]:
- """Generate elements from an async iterable,
- and a boolean indicating if it is the last element."""
- iterator = iterable.__aiter__()
- try:
- item = await iterator.__anext__()
- except StopAsyncIteration:
- return
- async for next_item in iterator:
- yield item, False
- item = next_item
- yield item, True
-
-
-class Tongyi(BaseLLM):
- """Tongyi completion model integration.
-
- Setup:
- Install ``dashscope`` and set environment variables ``DASHSCOPE_API_KEY``.
-
- .. code-block:: bash
-
- pip install dashscope
- export DASHSCOPE_API_KEY="your-api-key"
-
- Key init args — completion params:
- model: str
- Name of Tongyi model to use.
- top_p: float
- Total probability mass of tokens to consider at each step.
- streaming: bool
- Whether to stream the results or not.
-
- Key init args — client params:
- api_key: Optional[str]
- Dashscope API KEY. If not passed in will be read from env var DASHSCOPE_API_KEY.
- max_retries: int
- Maximum number of retries to make when generating.
-
- See full list of supported init args and their descriptions in the params section.
-
- Instantiate:
- .. code-block:: python
-
- from langchain_community.llms import Tongyi
-
- llm = Tongyi(
- model="qwen-max",
- # top_p="...",
- # api_key="...",
- # other params...
- )
-
- Invoke:
- .. code-block:: python
-
- input_text = "用50个字左右阐述,生命的意义在于"
- llm.invoke(input_text)
-
- .. code-block:: python
-
- '探索、成长、连接与爱——在有限的时间里,不断学习、体验、贡献并寻找与世界和谐共存之道,让每一刻充满价值与意义。'
-
- Stream:
- .. code-block:: python
-
- for chunk in llm.stream(input_text):
- print(chunk)
-
- .. code-block:: python
-
- 探索 | 、 | 成长 | 、连接与爱。 | 在有限的时间里,寻找个人价值, | 贡献于他人,共同体验世界的美好 | ,让世界因自己的存在而更 | 温暖。
-
- Async:
- .. code-block:: python
-
- await llm.ainvoke(input_text)
-
- # stream:
- # async for chunk in llm.astream(input_text):
- # print(chunk)
-
- # batch:
- # await llm.abatch([input_text])
-
- .. code-block:: python
-
- '探索、成长、连接与爱。在有限的时间里,寻找个人价值,贡献于他人和社会,体验丰富多彩的情感与经历,不断学习进步,让世界因自己的存在而更美好。'
-
- """ # noqa: E501
-
- @property
- def lc_secrets(self) -> Dict[str, str]:
- return {"dashscope_api_key": "DASHSCOPE_API_KEY"}
-
- client: Any = None #: :meta private:
- model_name: str = Field(default="qwen-plus", alias="model")
-
- """Model name to use."""
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
-
- top_p: float = 0.8
- """Total probability mass of tokens to consider at each step."""
-
- dashscope_api_key: Optional[str] = Field(default=None, alias="api_key")
- """Dashscope api key provide by Alibaba Cloud."""
-
- streaming: bool = False
- """Whether to stream the results or not."""
-
- max_retries: int = 10
- """Maximum number of retries to make when generating."""
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "tongyi"
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- values["dashscope_api_key"] = get_from_dict_or_env(
- values, ["dashscope_api_key", "api_key"], "DASHSCOPE_API_KEY"
- )
- try:
- import dashscope
- except ImportError:
- raise ImportError(
- "Could not import dashscope python package. "
- "Please install it with `pip install dashscope`."
- )
- try:
- values["client"] = dashscope.Generation
- except AttributeError:
- raise ValueError(
- "`dashscope` has no `Generation` attribute, this is likely "
- "due to an old version of the dashscope package. Try upgrading it "
- "with `pip install --upgrade dashscope`."
- )
-
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling Tongyi Qwen API."""
- normal_params = {
- "model": self.model_name,
- "top_p": self.top_p,
- "api_key": self.dashscope_api_key,
- }
-
- return {**normal_params, **self.model_kwargs}
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- return {"model_name": self.model_name, **super()._identifying_params}
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- generations = []
- if self.streaming:
- if len(prompts) > 1:
- raise ValueError("Cannot stream results with multiple prompts.")
- generation: Optional[GenerationChunk] = None
- for chunk in self._stream(prompts[0], stop, run_manager, **kwargs):
- if generation is None:
- generation = chunk
- else:
- generation += chunk
- assert generation is not None
- generations.append([self._chunk_to_generation(generation)])
- else:
- params: Dict[str, Any] = self._invocation_params(stop=stop, **kwargs)
- for prompt in prompts:
- completion = generate_with_retry(self, prompt=prompt, **params)
- generations.append(
- [Generation(**self._generation_from_qwen_resp(completion))]
- )
- return LLMResult(
- generations=generations,
- llm_output={
- "model_name": self.model_name,
- },
- )
-
- async def _agenerate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- generations = []
- if self.streaming:
- if len(prompts) > 1:
- raise ValueError("Cannot stream results with multiple prompts.")
- generation: Optional[GenerationChunk] = None
- async for chunk in self._astream(prompts[0], stop, run_manager, **kwargs):
- if generation is None:
- generation = chunk
- else:
- generation += chunk
- assert generation is not None
- generations.append([self._chunk_to_generation(generation)])
- else:
- params: Dict[str, Any] = self._invocation_params(stop=stop, **kwargs)
- for prompt in prompts:
- completion = await asyncio.get_running_loop().run_in_executor(
- None,
- functools.partial(
- generate_with_retry, **{"llm": self, "prompt": prompt, **params}
- ),
- )
- generations.append(
- [Generation(**self._generation_from_qwen_resp(completion))]
- )
- return LLMResult(
- generations=generations,
- llm_output={
- "model_name": self.model_name,
- },
- )
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- params: Dict[str, Any] = self._invocation_params(
- stop=stop, stream=True, **kwargs
- )
- for stream_resp, is_last_chunk in generate_with_last_element_mark(
- stream_generate_with_retry(self, prompt=prompt, **params)
- ):
- chunk = GenerationChunk(
- **self._generation_from_qwen_resp(stream_resp, is_last_chunk)
- )
- if run_manager:
- run_manager.on_llm_new_token(
- chunk.text,
- chunk=chunk,
- verbose=self.verbose,
- )
- yield chunk
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- params: Dict[str, Any] = self._invocation_params(
- stop=stop, stream=True, **kwargs
- )
- async for stream_resp, is_last_chunk in agenerate_with_last_element_mark(
- astream_generate_with_retry(self, prompt=prompt, **params)
- ):
- chunk = GenerationChunk(
- **self._generation_from_qwen_resp(stream_resp, is_last_chunk)
- )
- if run_manager:
- await run_manager.on_llm_new_token(
- chunk.text,
- chunk=chunk,
- verbose=self.verbose,
- )
- yield chunk
-
- def _invocation_params(self, stop: Any, **kwargs: Any) -> Dict[str, Any]:
- params = {
- **self._default_params,
- **kwargs,
- }
- if stop is not None:
- params["stop"] = stop
- if params.get("stream"):
- params["incremental_output"] = True
- return params
-
- @staticmethod
- def _generation_from_qwen_resp(
- resp: Any, is_last_chunk: bool = True
- ) -> Dict[str, Any]:
- # According to the response from dashscope,
- # each chunk's `generation_info` overwrites the previous one.
- # Besides, The `merge_dicts` method,
- # which is used to concatenate `generation_info` in `GenerationChunk`,
- # does not support merging of int type values.
- # Therefore, we adopt the `generation_info` of the last chunk
- # and discard the `generation_info` of the intermediate chunks.
- if is_last_chunk:
- return dict(
- text=resp["output"]["text"],
- generation_info=dict(
- finish_reason=resp["output"]["finish_reason"],
- request_id=resp["request_id"],
- token_usage=dict(resp["usage"]),
- ),
- )
- else:
- return dict(text=resp["output"]["text"])
-
- @staticmethod
- def _chunk_to_generation(chunk: GenerationChunk) -> Generation:
- return Generation(
- text=chunk.text,
- generation_info=chunk.generation_info,
- )
diff --git a/libs/community/langchain_community/llms/utils.py b/libs/community/langchain_community/llms/utils.py
deleted file mode 100644
index 15c30d59c3..0000000000
--- a/libs/community/langchain_community/llms/utils.py
+++ /dev/null
@@ -1,9 +0,0 @@
-"""Common utility functions for LLM APIs."""
-
-import re
-from typing import List
-
-
-def enforce_stop_tokens(text: str, stop: List[str]) -> str:
- """Cut off the text as soon as any stop words occur."""
- return re.split("|".join(stop), text, maxsplit=1)[0]
diff --git a/libs/community/langchain_community/llms/vertexai.py b/libs/community/langchain_community/llms/vertexai.py
deleted file mode 100644
index 74ec9374ac..0000000000
--- a/libs/community/langchain_community/llms/vertexai.py
+++ /dev/null
@@ -1,542 +0,0 @@
-from __future__ import annotations
-
-from concurrent.futures import Executor, ThreadPoolExecutor
-from typing import TYPE_CHECKING, Any, ClassVar, Dict, Iterator, List, Optional, Union
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks.manager import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import BaseLLM
-from langchain_core.outputs import Generation, GenerationChunk, LLMResult
-from langchain_core.utils import pre_init
-from pydantic import BaseModel, ConfigDict, Field
-
-from langchain_community.utilities.vertexai import (
- create_retry_decorator,
- get_client_info,
- init_vertexai,
- raise_vertex_import_error,
-)
-
-if TYPE_CHECKING:
- from google.cloud.aiplatform.gapic import (
- PredictionServiceAsyncClient,
- PredictionServiceClient,
- )
- from google.cloud.aiplatform.models import Prediction
- from google.protobuf.struct_pb2 import Value
- from vertexai.language_models._language_models import (
- TextGenerationResponse,
- _LanguageModel,
- )
- from vertexai.preview.generative_models import Image
-
-# This is for backwards compatibility
-# We can remove after `langchain` stops importing it
-_response_to_generation = None
-stream_completion_with_retry = None
-
-
-def is_codey_model(model_name: str) -> bool:
- """Return True if the model name is a Codey model."""
- return "code" in model_name
-
-
-def is_gemini_model(model_name: str) -> bool:
- """Return True if the model name is a Gemini model."""
- return model_name is not None and "gemini" in model_name
-
-
-def completion_with_retry(
- llm: VertexAI,
- prompt: List[Union[str, "Image"]],
- stream: bool = False,
- is_gemini: bool = False,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
-) -> Any:
- """Use tenacity to retry the completion call."""
- retry_decorator = create_retry_decorator(llm, run_manager=run_manager)
-
- @retry_decorator
- def _completion_with_retry(
- prompt: List[Union[str, "Image"]], is_gemini: bool = False, **kwargs: Any
- ) -> Any:
- if is_gemini:
- return llm.client.generate_content(
- prompt, stream=stream, generation_config=kwargs
- )
- else:
- if stream:
- return llm.client.predict_streaming(prompt[0], **kwargs)
- return llm.client.predict(prompt[0], **kwargs)
-
- return _completion_with_retry(prompt, is_gemini, **kwargs)
-
-
-async def acompletion_with_retry(
- llm: VertexAI,
- prompt: str,
- is_gemini: bool = False,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
-) -> Any:
- """Use tenacity to retry the completion call."""
- retry_decorator = create_retry_decorator(llm, run_manager=run_manager)
-
- @retry_decorator
- async def _acompletion_with_retry(
- prompt: str, is_gemini: bool = False, **kwargs: Any
- ) -> Any:
- if is_gemini:
- return await llm.client.generate_content_async(
- prompt, generation_config=kwargs
- )
- return await llm.client.predict_async(prompt, **kwargs)
-
- return await _acompletion_with_retry(prompt, is_gemini, **kwargs)
-
-
-class _VertexAIBase(BaseModel):
- model_config = ConfigDict(protected_namespaces=())
-
- project: Optional[str] = None
- "The default GCP project to use when making Vertex API calls."
- location: str = "us-central1"
- "The default location to use when making API calls."
- request_parallelism: int = 5
- "The amount of parallelism allowed for requests issued to VertexAI models. "
- "Default is 5."
- max_retries: int = 6
- """The maximum number of retries to make when generating."""
- task_executor: ClassVar[Optional[Executor]] = Field(default=None, exclude=True)
- stop: Optional[List[str]] = None
- "Optional list of stop words to use when generating."
- model_name: Optional[str] = None
- "Underlying model name."
-
- @classmethod
- def _get_task_executor(cls, request_parallelism: int = 5) -> Executor:
- if cls.task_executor is None:
- cls.task_executor = ThreadPoolExecutor(max_workers=request_parallelism)
- return cls.task_executor
-
-
-class _VertexAICommon(_VertexAIBase):
- client: "_LanguageModel" = None #: :meta private:
- client_preview: "_LanguageModel" = None #: :meta private:
- model_name: str
- "Underlying model name."
- temperature: float = 0.0
- "Sampling temperature, it controls the degree of randomness in token selection."
- max_output_tokens: int = 128
- "Token limit determines the maximum amount of text output from one prompt."
- top_p: float = 0.95
- "Tokens are selected from most probable to least until the sum of their "
- "probabilities equals the top-p value. Top-p is ignored for Codey models."
- top_k: int = 40
- "How the model selects tokens for output, the next token is selected from "
- "among the top-k most probable tokens. Top-k is ignored for Codey models."
- credentials: Any = Field(default=None, exclude=True)
- "The default custom credentials (google.auth.credentials.Credentials) to use "
- "when making API calls. If not provided, credentials will be ascertained from "
- "the environment."
- n: int = 1
- """How many completions to generate for each prompt."""
- streaming: bool = False
- """Whether to stream the results or not."""
-
- @property
- def _llm_type(self) -> str:
- return "vertexai"
-
- @property
- def is_codey_model(self) -> bool:
- return is_codey_model(self.model_name)
-
- @property
- def _is_gemini_model(self) -> bool:
- return is_gemini_model(self.model_name)
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Gets the identifying parameters."""
- return {**{"model_name": self.model_name}, **self._default_params}
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- params = {
- "temperature": self.temperature,
- "max_output_tokens": self.max_output_tokens,
- "candidate_count": self.n,
- }
- if not self.is_codey_model:
- params.update(
- {
- "top_k": self.top_k,
- "top_p": self.top_p,
- }
- )
- return params
-
- @classmethod
- def _try_init_vertexai(cls, values: Dict) -> None:
- allowed_params = ["project", "location", "credentials"]
- params = {k: v for k, v in values.items() if k in allowed_params}
- init_vertexai(**params)
- return None
-
- def _prepare_params(
- self,
- stop: Optional[List[str]] = None,
- stream: bool = False,
- **kwargs: Any,
- ) -> dict:
- stop_sequences = stop or self.stop
- params_mapping = {"n": "candidate_count"}
- params = {params_mapping.get(k, k): v for k, v in kwargs.items()}
- params = {**self._default_params, "stop_sequences": stop_sequences, **params}
- if stream or self.streaming:
- params.pop("candidate_count")
- return params
-
-
-@deprecated(
- since="0.0.12",
- removal="1.0",
- alternative_import="langchain_google_vertexai.VertexAI",
-)
-class VertexAI(_VertexAICommon, BaseLLM):
- """Google Vertex AI large language models."""
-
- model_name: str = "text-bison"
- "The name of the Vertex AI large language model."
- tuned_model_name: Optional[str] = None
- "The name of a tuned model. If provided, model_name is ignored."
-
- @classmethod
- def is_lc_serializable(self) -> bool:
- return True
-
- @classmethod
- def get_lc_namespace(cls) -> List[str]:
- """Get the namespace of the langchain object."""
- return ["langchain", "llms", "vertexai"]
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that the python package exists in environment."""
- tuned_model_name = values.get("tuned_model_name")
- model_name = values["model_name"]
- is_gemini = is_gemini_model(values["model_name"])
- cls._try_init_vertexai(values)
- try:
- from vertexai.language_models import (
- CodeGenerationModel,
- TextGenerationModel,
- )
- from vertexai.preview.language_models import (
- CodeGenerationModel as PreviewCodeGenerationModel,
- )
- from vertexai.preview.language_models import (
- TextGenerationModel as PreviewTextGenerationModel,
- )
-
- if is_gemini:
- from vertexai.preview.generative_models import (
- GenerativeModel,
- )
-
- if is_codey_model(model_name):
- model_cls = CodeGenerationModel
- preview_model_cls = PreviewCodeGenerationModel
- elif is_gemini:
- model_cls = GenerativeModel
- preview_model_cls = GenerativeModel
- else:
- model_cls = TextGenerationModel
- preview_model_cls = PreviewTextGenerationModel
-
- if tuned_model_name:
- values["client"] = model_cls.get_tuned_model(tuned_model_name)
- values["client_preview"] = preview_model_cls.get_tuned_model(
- tuned_model_name
- )
- else:
- if is_gemini:
- values["client"] = model_cls(model_name=model_name)
- values["client_preview"] = preview_model_cls(model_name=model_name)
- else:
- values["client"] = model_cls.from_pretrained(model_name)
- values["client_preview"] = preview_model_cls.from_pretrained(
- model_name
- )
-
- except ImportError:
- raise_vertex_import_error()
-
- if values["streaming"] and values["n"] > 1:
- raise ValueError("Only one candidate can be generated with streaming!")
- return values
-
- def get_num_tokens(self, text: str) -> int:
- """Get the number of tokens present in the text.
-
- Useful for checking if an input will fit in a model's context window.
-
- Args:
- text: The string input to tokenize.
-
- Returns:
- The integer number of tokens in the text.
- """
- try:
- result = self.client_preview.count_tokens([text])
- except AttributeError:
- raise_vertex_import_error()
-
- return result.total_tokens
-
- def _response_to_generation(
- self, response: TextGenerationResponse
- ) -> GenerationChunk:
- """Converts a stream response to a generation chunk."""
- try:
- generation_info = {
- "is_blocked": response.is_blocked,
- "safety_attributes": response.safety_attributes,
- }
- except Exception:
- generation_info = None
- return GenerationChunk(text=response.text, generation_info=generation_info)
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- stream: Optional[bool] = None,
- **kwargs: Any,
- ) -> LLMResult:
- should_stream = stream if stream is not None else self.streaming
- params = self._prepare_params(stop=stop, stream=should_stream, **kwargs)
- generations: List[List[Generation]] = []
- for prompt in prompts:
- if should_stream:
- generation = GenerationChunk(text="")
- for chunk in self._stream(
- prompt, stop=stop, run_manager=run_manager, **kwargs
- ):
- generation += chunk
- generations.append([generation])
- else:
- res = completion_with_retry(
- self,
- [prompt],
- stream=should_stream,
- is_gemini=self._is_gemini_model,
- run_manager=run_manager,
- **params,
- )
- generations.append(
- [self._response_to_generation(r) for r in res.candidates]
- )
- return LLMResult(generations=generations)
-
- async def _agenerate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- params = self._prepare_params(stop=stop, **kwargs)
- generations = []
- for prompt in prompts:
- res = await acompletion_with_retry(
- self,
- prompt,
- is_gemini=self._is_gemini_model,
- run_manager=run_manager,
- **params,
- )
- generations.append(
- [self._response_to_generation(r) for r in res.candidates]
- )
- return LLMResult(generations=generations) # type: ignore[arg-type]
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- params = self._prepare_params(stop=stop, stream=True, **kwargs)
- for stream_resp in completion_with_retry(
- self,
- [prompt],
- stream=True,
- is_gemini=self._is_gemini_model,
- run_manager=run_manager,
- **params,
- ):
- chunk = self._response_to_generation(stream_resp)
- if run_manager:
- run_manager.on_llm_new_token(
- chunk.text,
- chunk=chunk,
- verbose=self.verbose,
- )
- yield chunk
-
-
-@deprecated(
- since="0.0.12",
- removal="1.0",
- alternative_import="langchain_google_vertexai.VertexAIModelGarden",
-)
-class VertexAIModelGarden(_VertexAIBase, BaseLLM):
- """Vertex AI Model Garden large language models."""
-
- client: "PredictionServiceClient" = (
- None #: :meta private: # type: ignore[assignment]
- )
- async_client: "PredictionServiceAsyncClient" = (
- None #: :meta private: # type: ignore[assignment]
- )
- endpoint_id: str
- "A name of an endpoint where the model has been deployed."
- allowed_model_args: Optional[List[str]] = None
- "Allowed optional args to be passed to the model."
- prompt_arg: str = "prompt"
- result_arg: Optional[str] = "generated_text"
- "Set result_arg to None if output of the model is expected to be a string."
- "Otherwise, if it's a dict, provided an argument that contains the result."
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that the python package exists in environment."""
- try:
- from google.api_core.client_options import ClientOptions
- from google.cloud.aiplatform.gapic import (
- PredictionServiceAsyncClient,
- PredictionServiceClient,
- )
- except ImportError:
- raise_vertex_import_error()
-
- if not values["project"]:
- raise ValueError(
- "A GCP project should be provided to run inference on Model Garden!"
- )
-
- client_options = ClientOptions(
- api_endpoint=f"{values['location']}-aiplatform.googleapis.com"
- )
- client_info = get_client_info(module="vertex-ai-model-garden")
- values["client"] = PredictionServiceClient(
- client_options=client_options, client_info=client_info
- )
- values["async_client"] = PredictionServiceAsyncClient(
- client_options=client_options, client_info=client_info
- )
- return values
-
- @property
- def endpoint_path(self) -> str:
- return self.client.endpoint_path(
- project=self.project,
- location=self.location,
- endpoint=self.endpoint_id,
- )
-
- @property
- def _llm_type(self) -> str:
- return "vertexai_model_garden"
-
- def _prepare_request(self, prompts: List[str], **kwargs: Any) -> List["Value"]:
- try:
- from google.protobuf import json_format
- from google.protobuf.struct_pb2 import Value
- except ImportError:
- raise ImportError(
- "protobuf package not found, please install it with"
- " `pip install protobuf`"
- )
- instances = []
- for prompt in prompts:
- if self.allowed_model_args:
- instance = {
- k: v for k, v in kwargs.items() if k in self.allowed_model_args
- }
- else:
- instance = {}
- instance[self.prompt_arg] = prompt
- instances.append(instance)
-
- predict_instances = [
- json_format.ParseDict(instance_dict, Value()) for instance_dict in instances
- ]
- return predict_instances
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Run the LLM on the given prompt and input."""
- instances = self._prepare_request(prompts, **kwargs)
- response = self.client.predict(endpoint=self.endpoint_path, instances=instances)
- return self._parse_response(response)
-
- def _parse_response(self, predictions: "Prediction") -> LLMResult:
- generations: List[List[Generation]] = []
- for result in predictions.predictions:
- generations.append(
- [
- Generation(text=self._parse_prediction(prediction))
- for prediction in result
- ]
- )
- return LLMResult(generations=generations)
-
- def _parse_prediction(self, prediction: Any) -> str:
- if isinstance(prediction, str):
- return prediction
-
- if self.result_arg:
- try:
- return prediction[self.result_arg]
- except KeyError:
- if isinstance(prediction, str):
- error_desc = (
- "Provided non-None `result_arg` (result_arg="
- f"{self.result_arg}). But got prediction of type "
- f"{type(prediction)} instead of dict. Most probably, you"
- "need to set `result_arg=None` during VertexAIModelGarden "
- "initialization."
- )
- raise ValueError(error_desc)
- else:
- raise ValueError(f"{self.result_arg} key not found in prediction!")
-
- return prediction
-
- async def _agenerate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Run the LLM on the given prompt and input."""
- instances = self._prepare_request(prompts, **kwargs)
- response = await self.async_client.predict(
- endpoint=self.endpoint_path, instances=instances
- )
- return self._parse_response(response)
diff --git a/libs/community/langchain_community/llms/vllm.py b/libs/community/langchain_community/llms/vllm.py
deleted file mode 100644
index 66a0f17756..0000000000
--- a/libs/community/langchain_community/llms/vllm.py
+++ /dev/null
@@ -1,189 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import BaseLLM
-from langchain_core.outputs import Generation, LLMResult
-from langchain_core.utils import pre_init
-from pydantic import Field
-
-from langchain_community.llms.openai import BaseOpenAI
-from langchain_community.utils.openai import is_openai_v1
-
-
-class VLLM(BaseLLM):
- """VLLM language model."""
-
- model: str = ""
- """The name or path of a HuggingFace Transformers model."""
-
- tensor_parallel_size: Optional[int] = 1
- """The number of GPUs to use for distributed execution with tensor parallelism."""
-
- trust_remote_code: Optional[bool] = False
- """Trust remote code (e.g., from HuggingFace) when downloading the model
- and tokenizer."""
-
- n: int = 1
- """Number of output sequences to return for the given prompt."""
-
- best_of: Optional[int] = None
- """Number of output sequences that are generated from the prompt."""
-
- presence_penalty: float = 0.0
- """Float that penalizes new tokens based on whether they appear in the
- generated text so far"""
-
- frequency_penalty: float = 0.0
- """Float that penalizes new tokens based on their frequency in the
- generated text so far"""
-
- temperature: float = 1.0
- """Float that controls the randomness of the sampling."""
-
- top_p: float = 1.0
- """Float that controls the cumulative probability of the top tokens to consider."""
-
- top_k: int = -1
- """Integer that controls the number of top tokens to consider."""
-
- use_beam_search: bool = False
- """Whether to use beam search instead of sampling."""
-
- stop: Optional[List[str]] = None
- """List of strings that stop the generation when they are generated."""
-
- ignore_eos: bool = False
- """Whether to ignore the EOS token and continue generating tokens after
- the EOS token is generated."""
-
- max_new_tokens: int = 512
- """Maximum number of tokens to generate per output sequence."""
-
- logprobs: Optional[int] = None
- """Number of log probabilities to return per output token."""
-
- dtype: str = "auto"
- """The data type for the model weights and activations."""
-
- download_dir: Optional[str] = None
- """Directory to download and load the weights. (Default to the default
- cache dir of huggingface)"""
-
- vllm_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `vllm.LLM` call not explicitly specified."""
-
- client: Any = None #: :meta private:
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that python package exists in environment."""
-
- try:
- from vllm import LLM as VLLModel
- except ImportError:
- raise ImportError(
- "Could not import vllm python package. "
- "Please install it with `pip install vllm`."
- )
-
- values["client"] = VLLModel(
- model=values["model"],
- tensor_parallel_size=values["tensor_parallel_size"],
- trust_remote_code=values["trust_remote_code"],
- dtype=values["dtype"],
- download_dir=values["download_dir"],
- **values["vllm_kwargs"],
- )
-
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling vllm."""
- return {
- "n": self.n,
- "best_of": self.best_of,
- "max_tokens": self.max_new_tokens,
- "top_k": self.top_k,
- "top_p": self.top_p,
- "temperature": self.temperature,
- "presence_penalty": self.presence_penalty,
- "frequency_penalty": self.frequency_penalty,
- "stop": self.stop,
- "ignore_eos": self.ignore_eos,
- "use_beam_search": self.use_beam_search,
- "logprobs": self.logprobs,
- }
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Run the LLM on the given prompt and input."""
- from vllm import SamplingParams
-
- lora_request = kwargs.pop("lora_request", None)
-
- # build sampling parameters
- params = {**self._default_params, **kwargs, "stop": stop}
-
- # filter params for SamplingParams
- known_keys = SamplingParams.__annotations__.keys()
- sample_params = SamplingParams(
- **{k: v for k, v in params.items() if k in known_keys}
- )
-
- # call the model
- if lora_request:
- outputs = self.client.generate(
- prompts, sample_params, lora_request=lora_request
- )
- else:
- outputs = self.client.generate(prompts, sample_params)
-
- generations = []
- for output in outputs:
- text = output.outputs[0].text
- generations.append([Generation(text=text)])
-
- return LLMResult(generations=generations)
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "vllm"
-
-
-class VLLMOpenAI(BaseOpenAI):
- """vLLM OpenAI-compatible API client"""
-
- @classmethod
- def is_lc_serializable(cls) -> bool:
- return False
-
- @property
- def _invocation_params(self) -> Dict[str, Any]:
- """Get the parameters used to invoke the model."""
-
- params: Dict[str, Any] = {
- "model": self.model_name,
- **self._default_params,
- "logit_bias": None,
- }
- if not is_openai_v1():
- params.update(
- {
- "api_key": self.openai_api_key,
- "api_base": self.openai_api_base,
- }
- )
-
- return params
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "vllm-openai"
diff --git a/libs/community/langchain_community/llms/volcengine_maas.py b/libs/community/langchain_community/llms/volcengine_maas.py
deleted file mode 100644
index e737a0a556..0000000000
--- a/libs/community/langchain_community/llms/volcengine_maas.py
+++ /dev/null
@@ -1,182 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Dict, Iterator, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import BaseModel, ConfigDict, Field, SecretStr
-
-
-class VolcEngineMaasBase(BaseModel):
- """Base class for VolcEngineMaas models."""
-
- model_config = ConfigDict(protected_namespaces=())
-
- client: Any = None
-
- volc_engine_maas_ak: Optional[SecretStr] = None
- """access key for volc engine"""
- volc_engine_maas_sk: Optional[SecretStr] = None
- """secret key for volc engine"""
-
- endpoint: Optional[str] = "maas-api.ml-platform-cn-beijing.volces.com"
- """Endpoint of the VolcEngineMaas LLM."""
-
- region: Optional[str] = "Region"
- """Region of the VolcEngineMaas LLM."""
-
- model: str = "skylark-lite-public"
- """Model name. you could check this model details here
- https://www.volcengine.com/docs/82379/1133187
- and you could choose other models by change this field"""
- model_version: Optional[str] = None
- """Model version. Only used in moonshot large language model.
- you could check details here https://www.volcengine.com/docs/82379/1158281"""
-
- top_p: Optional[float] = 0.8
- """Total probability mass of tokens to consider at each step."""
-
- temperature: Optional[float] = 0.95
- """A non-negative float that tunes the degree of randomness in generation."""
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """model special arguments, you could check detail on model page"""
-
- streaming: bool = False
- """Whether to stream the results."""
-
- connect_timeout: Optional[int] = 60
- """Timeout for connect to volc engine maas endpoint. Default is 60 seconds."""
-
- read_timeout: Optional[int] = 60
- """Timeout for read response from volc engine maas endpoint.
- Default is 60 seconds."""
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- volc_engine_maas_ak = convert_to_secret_str(
- get_from_dict_or_env(values, "volc_engine_maas_ak", "VOLC_ACCESSKEY")
- )
- volc_engine_maas_sk = convert_to_secret_str(
- get_from_dict_or_env(values, "volc_engine_maas_sk", "VOLC_SECRETKEY")
- )
- endpoint = values["endpoint"]
- if values["endpoint"] is not None and values["endpoint"] != "":
- endpoint = values["endpoint"]
- try:
- from volcengine.maas import MaasService
-
- maas = MaasService(
- endpoint,
- values["region"],
- connection_timeout=values["connect_timeout"],
- socket_timeout=values["read_timeout"],
- )
- maas.set_ak(volc_engine_maas_ak.get_secret_value())
- maas.set_sk(volc_engine_maas_sk.get_secret_value())
-
- values["volc_engine_maas_ak"] = volc_engine_maas_ak
- values["volc_engine_maas_sk"] = volc_engine_maas_sk
- values["client"] = maas
- except ImportError:
- raise ImportError(
- "volcengine package not found, please install it with "
- "`pip install volcengine`"
- )
- return values
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- """Get the default parameters for calling VolcEngineMaas API."""
- normal_params = {
- "top_p": self.top_p,
- "temperature": self.temperature,
- }
-
- return {**normal_params, **self.model_kwargs}
-
-
-class VolcEngineMaasLLM(LLM, VolcEngineMaasBase):
- """volc engine maas hosts a plethora of models.
- You can utilize these models through this class.
-
- To use, you should have the ``volcengine`` python package installed.
- and set access key and secret key by environment variable or direct pass those to
- this class.
- access key, secret key are required parameters which you could get help
- https://www.volcengine.com/docs/6291/65568
-
- In order to use them, it is necessary to install the 'volcengine' Python package.
- The access key and secret key must be set either via environment variables or
- passed directly to this class.
- access key and secret key are mandatory parameters for which assistance can be
- sought at https://www.volcengine.com/docs/6291/65568.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import VolcEngineMaasLLM
- model = VolcEngineMaasLLM(model="skylark-lite-public",
- volc_engine_maas_ak="your_ak",
- volc_engine_maas_sk="your_sk")
- """
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "volc-engine-maas-llm"
-
- def _convert_prompt_msg_params(
- self,
- prompt: str,
- **kwargs: Any,
- ) -> dict:
- model_req = {
- "model": {
- "name": self.model,
- }
- }
- if self.model_version is not None:
- model_req["model"]["version"] = self.model_version
-
- return {
- **model_req,
- "messages": [{"role": "user", "content": prompt}],
- "parameters": {**self._default_params, **kwargs},
- }
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- if self.streaming:
- completion = ""
- for chunk in self._stream(prompt, stop, run_manager, **kwargs):
- completion += chunk.text
- return completion
- params = self._convert_prompt_msg_params(prompt, **kwargs)
- response = self.client.chat(params)
-
- return response.get("choice", {}).get("message", {}).get("content", "")
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- params = self._convert_prompt_msg_params(prompt, **kwargs)
- for res in self.client.stream_chat(params):
- if res:
- chunk = GenerationChunk(
- text=res.get("choice", {}).get("message", {}).get("content", "")
- )
- if run_manager:
- run_manager.on_llm_new_token(chunk.text, chunk=chunk)
- yield chunk
diff --git a/libs/community/langchain_community/llms/watsonxllm.py b/libs/community/langchain_community/llms/watsonxllm.py
deleted file mode 100644
index 9a63d82413..0000000000
--- a/libs/community/langchain_community/llms/watsonxllm.py
+++ /dev/null
@@ -1,403 +0,0 @@
-import logging
-import os
-from typing import Any, Dict, Iterator, List, Mapping, Optional, Union
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import BaseLLM
-from langchain_core.outputs import Generation, GenerationChunk, LLMResult
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import ConfigDict, SecretStr
-
-logger = logging.getLogger(__name__)
-
-
-@deprecated(
- since="0.0.18", removal="1.0", alternative_import="langchain_ibm.WatsonxLLM"
-)
-class WatsonxLLM(BaseLLM):
- """
- IBM watsonx.ai large language models.
-
- To use, you should have ``ibm_watsonx_ai`` python package installed,
- and the environment variable ``WATSONX_APIKEY`` set with your API key, or pass
- it as a named parameter to the constructor.
-
-
- Example:
- .. code-block:: python
-
- from ibm_watsonx_ai.metanames import GenTextParamsMetaNames
- parameters = {
- GenTextParamsMetaNames.DECODING_METHOD: "sample",
- GenTextParamsMetaNames.MAX_NEW_TOKENS: 100,
- GenTextParamsMetaNames.MIN_NEW_TOKENS: 1,
- GenTextParamsMetaNames.TEMPERATURE: 0.5,
- GenTextParamsMetaNames.TOP_K: 50,
- GenTextParamsMetaNames.TOP_P: 1,
- }
-
- from langchain_community.llms import WatsonxLLM
- watsonx_llm = WatsonxLLM(
- model_id="google/flan-ul2",
- url="https://us-south.ml.cloud.ibm.com",
- apikey="*****",
- project_id="*****",
- params=parameters,
- )
- """
-
- model_id: str = ""
- """Type of model to use."""
-
- deployment_id: str = ""
- """Type of deployed model to use."""
-
- project_id: str = ""
- """ID of the Watson Studio project."""
-
- space_id: str = ""
- """ID of the Watson Studio space."""
-
- url: Optional[SecretStr] = None
- """Url to Watson Machine Learning instance"""
-
- apikey: Optional[SecretStr] = None
- """Apikey to Watson Machine Learning instance"""
-
- token: Optional[SecretStr] = None
- """Token to Watson Machine Learning instance"""
-
- password: Optional[SecretStr] = None
- """Password to Watson Machine Learning instance"""
-
- username: Optional[SecretStr] = None
- """Username to Watson Machine Learning instance"""
-
- instance_id: Optional[SecretStr] = None
- """Instance_id of Watson Machine Learning instance"""
-
- version: Optional[SecretStr] = None
- """Version of Watson Machine Learning instance"""
-
- params: Optional[dict] = None
- """Model parameters to use during generate requests."""
-
- verify: Union[str, bool] = ""
- """User can pass as verify one of following:
- the path to a CA_BUNDLE file
- the path of directory with certificates of trusted CAs
- True - default path to truststore will be taken
- False - no verification will be made"""
-
- streaming: bool = False
- """ Whether to stream the results or not. """
-
- watsonx_model: Any = None
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @classmethod
- def is_lc_serializable(cls) -> bool:
- return False
-
- @property
- def lc_secrets(self) -> Dict[str, str]:
- return {
- "url": "WATSONX_URL",
- "apikey": "WATSONX_APIKEY",
- "token": "WATSONX_TOKEN",
- "password": "WATSONX_PASSWORD",
- "username": "WATSONX_USERNAME",
- "instance_id": "WATSONX_INSTANCE_ID",
- }
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that credentials and python package exists in environment."""
- values["url"] = convert_to_secret_str(
- get_from_dict_or_env(values, "url", "WATSONX_URL")
- )
- if "cloud.ibm.com" in values.get("url", "").get_secret_value():
- values["apikey"] = convert_to_secret_str(
- get_from_dict_or_env(values, "apikey", "WATSONX_APIKEY")
- )
- else:
- if (
- not values["token"]
- and "WATSONX_TOKEN" not in os.environ
- and not values["password"]
- and "WATSONX_PASSWORD" not in os.environ
- and not values["apikey"]
- and "WATSONX_APIKEY" not in os.environ
- ):
- raise ValueError(
- "Did not find 'token', 'password' or 'apikey',"
- " please add an environment variable"
- " `WATSONX_TOKEN`, 'WATSONX_PASSWORD' or 'WATSONX_APIKEY' "
- "which contains it,"
- " or pass 'token', 'password' or 'apikey'"
- " as a named parameter."
- )
- elif values["token"] or "WATSONX_TOKEN" in os.environ:
- values["token"] = convert_to_secret_str(
- get_from_dict_or_env(values, "token", "WATSONX_TOKEN")
- )
- elif values["password"] or "WATSONX_PASSWORD" in os.environ:
- values["password"] = convert_to_secret_str(
- get_from_dict_or_env(values, "password", "WATSONX_PASSWORD")
- )
- values["username"] = convert_to_secret_str(
- get_from_dict_or_env(values, "username", "WATSONX_USERNAME")
- )
- elif values["apikey"] or "WATSONX_APIKEY" in os.environ:
- values["apikey"] = convert_to_secret_str(
- get_from_dict_or_env(values, "apikey", "WATSONX_APIKEY")
- )
- values["username"] = convert_to_secret_str(
- get_from_dict_or_env(values, "username", "WATSONX_USERNAME")
- )
- if not values["instance_id"] or "WATSONX_INSTANCE_ID" not in os.environ:
- values["instance_id"] = convert_to_secret_str(
- get_from_dict_or_env(values, "instance_id", "WATSONX_INSTANCE_ID")
- )
-
- try:
- from ibm_watsonx_ai.foundation_models import ModelInference
-
- credentials = {
- "url": values["url"].get_secret_value() if values["url"] else None,
- "apikey": (
- values["apikey"].get_secret_value() if values["apikey"] else None
- ),
- "token": (
- values["token"].get_secret_value() if values["token"] else None
- ),
- "password": (
- values["password"].get_secret_value()
- if values["password"]
- else None
- ),
- "username": (
- values["username"].get_secret_value()
- if values["username"]
- else None
- ),
- "instance_id": (
- values["instance_id"].get_secret_value()
- if values["instance_id"]
- else None
- ),
- "version": (
- values["version"].get_secret_value() if values["version"] else None
- ),
- }
- credentials_without_none_value = {
- key: value for key, value in credentials.items() if value is not None
- }
-
- watsonx_model = ModelInference(
- model_id=values["model_id"],
- deployment_id=values["deployment_id"],
- credentials=credentials_without_none_value,
- params=values["params"],
- project_id=values["project_id"],
- space_id=values["space_id"],
- verify=values["verify"],
- )
- values["watsonx_model"] = watsonx_model
-
- except ImportError:
- raise ImportError(
- "Could not import ibm_watsonx_ai python package. "
- "Please install it with `pip install ibm_watsonx_ai`."
- )
- return values
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- "model_id": self.model_id,
- "deployment_id": self.deployment_id,
- "params": self.params,
- "project_id": self.project_id,
- "space_id": self.space_id,
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "IBM watsonx.ai"
-
- @staticmethod
- def _extract_token_usage(
- response: Optional[List[Dict[str, Any]]] = None,
- ) -> Dict[str, Any]:
- if response is None:
- return {"generated_token_count": 0, "input_token_count": 0}
-
- input_token_count = 0
- generated_token_count = 0
-
- def get_count_value(key: str, result: Dict[str, Any]) -> int:
- return result.get(key, 0) or 0
-
- for res in response:
- results = res.get("results")
- if results:
- input_token_count += get_count_value("input_token_count", results[0])
- generated_token_count += get_count_value(
- "generated_token_count", results[0]
- )
-
- return {
- "generated_token_count": generated_token_count,
- "input_token_count": input_token_count,
- }
-
- def _get_chat_params(self, stop: Optional[List[str]] = None) -> Dict[str, Any]:
- params: Dict[str, Any] = {**self.params} if self.params else {}
- if stop is not None:
- params["stop_sequences"] = stop
- return params
-
- def _create_llm_result(self, response: List[dict]) -> LLMResult:
- """Create the LLMResult from the choices and prompts."""
- generations = []
- for res in response:
- results = res.get("results")
- if results:
- finish_reason = results[0].get("stop_reason")
- gen = Generation(
- text=results[0].get("generated_text"),
- generation_info={"finish_reason": finish_reason},
- )
- generations.append([gen])
- final_token_usage = self._extract_token_usage(response)
- llm_output = {
- "token_usage": final_token_usage,
- "model_id": self.model_id,
- "deployment_id": self.deployment_id,
- }
- return LLMResult(generations=generations, llm_output=llm_output)
-
- def _stream_response_to_generation_chunk(
- self,
- stream_response: Dict[str, Any],
- ) -> GenerationChunk:
- """Convert a stream response to a generation chunk."""
- if not stream_response["results"]:
- return GenerationChunk(text="")
- return GenerationChunk(
- text=stream_response["results"][0]["generated_text"],
- generation_info=dict(
- finish_reason=stream_response["results"][0].get("stop_reason", None),
- llm_output={
- "model_id": self.model_id,
- "deployment_id": self.deployment_id,
- },
- ),
- )
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call the IBM watsonx.ai inference endpoint.
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
- run_manager: Optional callback manager.
- Returns:
- The string generated by the model.
- Example:
- .. code-block:: python
-
- response = watsonx_llm.invoke("What is a molecule")
- """
- result = self._generate(
- prompts=[prompt], stop=stop, run_manager=run_manager, **kwargs
- )
- return result.generations[0][0].text
-
- def _generate(
- self,
- prompts: List[str],
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- stream: Optional[bool] = None,
- **kwargs: Any,
- ) -> LLMResult:
- """Call the IBM watsonx.ai inference endpoint which then generate the response.
- Args:
- prompts: List of strings (prompts) to pass into the model.
- stop: Optional list of stop words to use when generating.
- run_manager: Optional callback manager.
- Returns:
- The full LLMResult output.
- Example:
- .. code-block:: python
-
- response = watsonx_llm.generate(["What is a molecule"])
- """
- params = self._get_chat_params(stop=stop)
- should_stream = stream if stream is not None else self.streaming
- if should_stream:
- if len(prompts) > 1:
- raise ValueError(
- f"WatsonxLLM currently only supports single prompt, got {prompts}"
- )
- generation = GenerationChunk(text="")
- stream_iter = self._stream(
- prompts[0], stop=stop, run_manager=run_manager, **kwargs
- )
- for chunk in stream_iter:
- if generation is None:
- generation = chunk
- else:
- generation += chunk
- assert generation is not None
- if isinstance(generation.generation_info, dict):
- llm_output = generation.generation_info.pop("llm_output")
- return LLMResult(generations=[[generation]], llm_output=llm_output)
- return LLMResult(generations=[[generation]])
- else:
- response = self.watsonx_model.generate(prompt=prompts, params=params)
- return self._create_llm_result(response)
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- """Call the IBM watsonx.ai inference endpoint which then streams the response.
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
- run_manager: Optional callback manager.
- Returns:
- The iterator which yields generation chunks.
- Example:
- .. code-block:: python
-
- response = watsonx_llm.stream("What is a molecule")
- for chunk in response:
- print(chunk, end='') # noqa: T201
- """
- params = self._get_chat_params(stop=stop)
- for stream_resp in self.watsonx_model.generate_text_stream(
- prompt=prompt, raw_response=True, params=params
- ):
- chunk = self._stream_response_to_generation_chunk(stream_resp)
-
- if run_manager:
- run_manager.on_llm_new_token(chunk.text, chunk=chunk)
- yield chunk
diff --git a/libs/community/langchain_community/llms/weight_only_quantization.py b/libs/community/langchain_community/llms/weight_only_quantization.py
deleted file mode 100644
index 916734414f..0000000000
--- a/libs/community/langchain_community/llms/weight_only_quantization.py
+++ /dev/null
@@ -1,243 +0,0 @@
-import importlib
-from typing import Any, List, Mapping, Optional
-
-from langchain_core.callbacks.manager import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from pydantic import ConfigDict
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-DEFAULT_MODEL_ID = "google/flan-t5-large"
-DEFAULT_TASK = "text2text-generation"
-VALID_TASKS = ("text2text-generation", "text-generation", "summarization")
-
-
-class WeightOnlyQuantPipeline(LLM):
- """Weight only quantized model.
-
- To use, you should have the `intel-extension-for-transformers` packabge and
- `transformers` package installed.
- intel-extension-for-transformers:
- https://github.com/intel/intel-extension-for-transformers
-
- Example using from_model_id:
- .. code-block:: python
-
- from langchain_community.llms import WeightOnlyQuantPipeline
- from intel_extension_for_transformers.transformers import (
- WeightOnlyQuantConfig
- )
- config = WeightOnlyQuantConfig
- hf = WeightOnlyQuantPipeline.from_model_id(
- model_id="google/flan-t5-large",
- task="text2text-generation"
- pipeline_kwargs={"max_new_tokens": 10},
- quantization_config=config,
- )
- Example passing pipeline in directly:
- .. code-block:: python
-
- from langchain_community.llms import WeightOnlyQuantPipeline
- from intel_extension_for_transformers.transformers import (
- AutoModelForSeq2SeqLM
- )
- from intel_extension_for_transformers.transformers import (
- WeightOnlyQuantConfig
- )
- from transformers import AutoTokenizer, pipeline
-
- model_id = "google/flan-t5-large"
- tokenizer = AutoTokenizer.from_pretrained(model_id)
- config = WeightOnlyQuantConfig
- model = AutoModelForSeq2SeqLM.from_pretrained(
- model_id,
- quantization_config=config,
- )
- pipe = pipeline(
- "text-generation",
- model=model,
- tokenizer=tokenizer,
- max_new_tokens=10,
- )
- hf = WeightOnlyQuantPipeline(pipeline=pipe)
- """
-
- pipeline: Any = None #: :meta private:
- model_id: str = DEFAULT_MODEL_ID
- """Model name or local path to use."""
-
- model_kwargs: Optional[dict] = None
- """Key word arguments passed to the model."""
-
- pipeline_kwargs: Optional[dict] = None
- """Key word arguments passed to the pipeline."""
-
- model_config = ConfigDict(
- extra="allow",
- )
-
- @classmethod
- def from_model_id(
- cls,
- model_id: str,
- task: str,
- device: Optional[int] = -1,
- device_map: Optional[str] = None,
- model_kwargs: Optional[dict] = None,
- pipeline_kwargs: Optional[dict] = None,
- load_in_4bit: Optional[bool] = False,
- load_in_8bit: Optional[bool] = False,
- quantization_config: Optional[Any] = None,
- **kwargs: Any,
- ) -> LLM:
- """Construct the pipeline object from model_id and task."""
- if device_map is not None and (isinstance(device, int) and device > -1):
- raise ValueError("`Device` and `device_map` cannot be set simultaneously!")
- if importlib.util.find_spec("torch") is None:
- raise ValueError(
- "Weight only quantization pipeline only support PyTorch now!"
- )
-
- try:
- from intel_extension_for_transformers.transformers import (
- AutoModelForCausalLM,
- AutoModelForSeq2SeqLM,
- )
- from intel_extension_for_transformers.utils.utils import is_ipex_available
- from transformers import AutoTokenizer
- from transformers import pipeline as hf_pipeline
- except ImportError:
- raise ImportError(
- "Could not import transformers python package. "
- "Please install it with `pip install transformers` "
- "and `pip install intel-extension-for-transformers`."
- )
- if isinstance(device, int) and device >= 0:
- if not is_ipex_available():
- raise ValueError("Don't find out Intel GPU on this machine!")
- device_map = "xpu:" + str(device)
- elif isinstance(device, int) and device < 0:
- device = None
-
- if device is None:
- if device_map is None:
- device_map = "cpu"
-
- _model_kwargs = model_kwargs or {}
- tokenizer = AutoTokenizer.from_pretrained(model_id, **_model_kwargs)
-
- try:
- if task == "text-generation":
- model = AutoModelForCausalLM.from_pretrained(
- model_id,
- load_in_4bit=load_in_4bit,
- load_in_8bit=load_in_8bit,
- quantization_config=quantization_config,
- use_llm_runtime=False,
- device_map=device_map,
- **_model_kwargs,
- )
- elif task in ("text2text-generation", "summarization"):
- model = AutoModelForSeq2SeqLM.from_pretrained(
- model_id,
- load_in_4bit=load_in_4bit,
- load_in_8bit=load_in_8bit,
- quantization_config=quantization_config,
- use_llm_runtime=False,
- device_map=device_map,
- **_model_kwargs,
- )
- else:
- raise ValueError(
- f"Got invalid task {task}, "
- f"currently only {VALID_TASKS} are supported"
- )
- except ImportError as e:
- raise ImportError(
- f"Could not load the {task} model due to missing dependencies."
- ) from e
-
- if "trust_remote_code" in _model_kwargs:
- _model_kwargs = {
- k: v for k, v in _model_kwargs.items() if k != "trust_remote_code"
- }
- _pipeline_kwargs = pipeline_kwargs or {}
- pipeline = hf_pipeline(
- task=task,
- model=model,
- tokenizer=tokenizer,
- device=device,
- model_kwargs=_model_kwargs,
- **_pipeline_kwargs,
- )
- if pipeline.task not in VALID_TASKS:
- raise ValueError(
- f"Got invalid task {pipeline.task}, "
- f"currently only {VALID_TASKS} are supported"
- )
- return cls(
- pipeline=pipeline,
- model_id=model_id,
- model_kwargs=_model_kwargs,
- pipeline_kwargs=_pipeline_kwargs,
- **kwargs,
- )
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- "model_id": self.model_id,
- "model_kwargs": self.model_kwargs,
- "pipeline_kwargs": self.pipeline_kwargs,
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "weight_only_quantization"
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call the HuggingFace model and return the output.
-
- Args:
- prompt: The prompt to use for generation.
- stop: A list of strings to stop generation when encountered.
-
- Returns:
- The generated text.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import WeightOnlyQuantPipeline
- llm = WeightOnlyQuantPipeline.from_model_id(
- model_id="google/flan-t5-large",
- task="text2text-generation",
- )
- llm.invoke("This is a prompt.")
- """
- response = self.pipeline(prompt)
- if self.pipeline.task == "text-generation":
- # Text generation return includes the starter text.
- text = response[0]["generated_text"][len(prompt) :]
- elif self.pipeline.task == "text2text-generation":
- text = response[0]["generated_text"]
- elif self.pipeline.task == "summarization":
- text = response[0]["summary_text"]
- else:
- raise ValueError(
- f"Got invalid task {self.pipeline.task}, "
- f"currently only {VALID_TASKS} are supported"
- )
- if stop:
- # This is a bit hacky, but I can't figure out a better way to enforce
- # stop tokens when making calls to huggingface_hub.
- text = enforce_stop_tokens(text, stop)
- return text
diff --git a/libs/community/langchain_community/llms/writer.py b/libs/community/langchain_community/llms/writer.py
deleted file mode 100644
index e68909d06e..0000000000
--- a/libs/community/langchain_community/llms/writer.py
+++ /dev/null
@@ -1,197 +0,0 @@
-from typing import Any, AsyncIterator, Dict, Iterator, List, Mapping, Optional
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import ConfigDict, Field, SecretStr, model_validator
-
-
-class Writer(LLM):
- """Writer large language models.
-
- To use, you should have the ``writer-sdk`` Python package installed, and the
- environment variable ``WRITER_API_KEY`` set with your API key.
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import Writer as WriterLLM
- from writerai import Writer, AsyncWriter
-
- client = Writer()
- async_client = AsyncWriter()
-
- chat = WriterLLM(
- client=client,
- async_client=async_client
- )
- """
-
- client: Any = Field(default=None, exclude=True) #: :meta private:
- async_client: Any = Field(default=None, exclude=True) #: :meta private:
-
- api_key: Optional[SecretStr] = Field(default=None)
- """Writer API key."""
-
- model_name: str = Field(default="palmyra-x-003-instruct", alias="model")
- """Model name to use."""
-
- max_tokens: Optional[int] = None
- """The maximum number of tokens that the model can generate in the response."""
-
- temperature: Optional[float] = 0.7
- """Controls the randomness of the model's outputs. Higher values lead to more
- random outputs, while lower values make the model more deterministic."""
-
- top_p: Optional[float] = None
- """Used to control the nucleus sampling, where only the most probable tokens
- with a cumulative probability of top_p are considered for sampling, providing
- a way to fine-tune the randomness of predictions."""
-
- stop: Optional[List[str]] = None
- """Specifies stopping conditions for the model's output generation. This can
- be an array of strings or a single string that the model will look for as a
- signal to stop generating further tokens."""
-
- best_of: Optional[int] = None
- """Specifies the number of completions to generate and return the best one.
- Useful for generating multiple outputs and choosing the best based on some
- criteria."""
-
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
- """Holds any model parameters valid for `create` call not explicitly specified."""
-
- model_config = ConfigDict(populate_by_name=True)
-
- @property
- def _default_params(self) -> Mapping[str, Any]:
- """Get the default parameters for calling Writer API."""
- return {
- "max_tokens": self.max_tokens,
- "temperature": self.temperature,
- "top_p": self.top_p,
- "stop": self.stop,
- "best_of": self.best_of,
- **self.model_kwargs,
- }
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- "model": self.model_name,
- **self._default_params,
- }
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "writer"
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validates that api key is passed and creates Writer clients."""
- try:
- from writerai import AsyncClient, Client
- except ImportError as e:
- raise ImportError(
- "Could not import writerai python package. "
- "Please install it with `pip install writerai`."
- ) from e
-
- if not values.get("client"):
- values.update(
- {
- "client": Client(
- api_key=get_from_dict_or_env(
- values, "api_key", "WRITER_API_KEY"
- )
- )
- }
- )
-
- if not values.get("async_client"):
- values.update(
- {
- "async_client": AsyncClient(
- api_key=get_from_dict_or_env(
- values, "api_key", "WRITER_API_KEY"
- )
- )
- }
- )
-
- if not (
- type(values.get("client")) is Client
- and type(values.get("async_client")) is AsyncClient
- ):
- raise ValueError(
- "'client' attribute must be with type 'Client' and "
- "'async_client' must be with type 'AsyncClient' from 'writerai' package"
- )
-
- return values
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- params = {**self._identifying_params, **kwargs}
- if stop is not None:
- params.update({"stop": stop})
- text = self.client.completions.create(prompt=prompt, **params).choices[0].text
- return text
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[list[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- params = {**self._identifying_params, **kwargs}
- if stop is not None:
- params.update({"stop": stop})
- response = await self.async_client.completions.create(prompt=prompt, **params)
- text = response.choices[0].text
- return text
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[list[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- params = {**self._identifying_params, **kwargs, "stream": True}
- if stop is not None:
- params.update({"stop": stop})
- response = self.client.completions.create(prompt=prompt, **params)
- for chunk in response:
- if run_manager:
- run_manager.on_llm_new_token(chunk.value)
- yield GenerationChunk(text=chunk.value)
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[list[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- params = {**self._identifying_params, **kwargs, "stream": True}
- if stop is not None:
- params.update({"stop": stop})
- response = await self.async_client.completions.create(prompt=prompt, **params)
- async for chunk in response:
- if run_manager:
- await run_manager.on_llm_new_token(chunk.value)
- yield GenerationChunk(text=chunk.value)
diff --git a/libs/community/langchain_community/llms/xinference.py b/libs/community/langchain_community/llms/xinference.py
deleted file mode 100644
index f3f2aa734b..0000000000
--- a/libs/community/langchain_community/llms/xinference.py
+++ /dev/null
@@ -1,393 +0,0 @@
-from __future__ import annotations
-
-import json
-from typing import (
- TYPE_CHECKING,
- Any,
- AsyncIterator,
- Dict,
- Generator,
- Iterator,
- List,
- Mapping,
- Optional,
- Union,
-)
-
-import aiohttp
-import requests
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-
-if TYPE_CHECKING:
- from xinference.client import RESTfulChatModelHandle, RESTfulGenerateModelHandle
- from xinference.model.llm.core import LlamaCppGenerateConfig
-
-
-class Xinference(LLM):
- """`Xinference` large-scale model inference service.
-
- To use, you should have the xinference library installed:
-
- .. code-block:: bash
-
- pip install "xinference[all]"
-
- If you're simply using the services provided by Xinference, you can utilize the xinference_client package:
-
- .. code-block:: bash
-
- pip install xinference_client
-
- Check out: https://github.com/xorbitsai/inference
- To run, you need to start a Xinference supervisor on one server and Xinference workers on the other servers
-
- Example:
- To start a local instance of Xinference, run
-
- .. code-block:: bash
-
- $ xinference
-
- You can also deploy Xinference in a distributed cluster. Here are the steps:
-
- Starting the supervisor:
-
- .. code-block:: bash
-
- $ xinference-supervisor
-
- Starting the worker:
-
- .. code-block:: bash
-
- $ xinference-worker
-
- Then, launch a model using command line interface (CLI).
-
- Example:
-
- .. code-block:: bash
-
- $ xinference launch -n orca -s 3 -q q4_0
-
- It will return a model UID. Then, you can use Xinference with LangChain.
-
- Example:
-
- .. code-block:: python
-
- from langchain_community.llms import Xinference
-
- llm = Xinference(
- server_url="http://0.0.0.0:9997",
- model_uid = {model_uid} # replace model_uid with the model UID return from launching the model
- )
-
- llm.invoke(
- prompt="Q: where can we visit in the capital of France? A:",
- generate_config={"max_tokens": 1024, "stream": True},
- )
-
- Example:
-
- .. code-block:: python
-
- from langchain_community.llms import Xinference
- from langchain.prompts import PromptTemplate
-
- llm = Xinference(
- server_url="http://0.0.0.0:9997",
- model_uid={model_uid}, # replace model_uid with the model UID return from launching the model
- stream=True
- )
- prompt = PromptTemplate(
- input=['country'],
- template="Q: where can we visit in the capital of {country}? A:"
- )
- chain = prompt | llm
- chain.stream(input={'country': 'France'})
-
-
- To view all the supported builtin models, run:
-
- .. code-block:: bash
-
- $ xinference list --all
-
- """ # noqa: E501
-
- client: Optional[Any] = None
- server_url: Optional[str]
- """URL of the xinference server"""
- model_uid: Optional[str]
- """UID of the launched model"""
- model_kwargs: Dict[str, Any]
- """Keyword arguments to be passed to xinference.LLM"""
-
- def __init__(
- self,
- server_url: Optional[str] = None,
- model_uid: Optional[str] = None,
- api_key: Optional[str] = None,
- **model_kwargs: Any,
- ):
- try:
- from xinference.client import RESTfulClient
- except ImportError:
- try:
- from xinference_client import RESTfulClient
- except ImportError as e:
- raise ImportError(
- "Could not import RESTfulClient from xinference. Please install it"
- " with `pip install xinference` or `pip install xinference_client`."
- ) from e
-
- model_kwargs = model_kwargs or {}
-
- super().__init__(
- **{ # type: ignore[arg-type]
- "server_url": server_url,
- "model_uid": model_uid,
- "model_kwargs": model_kwargs,
- }
- )
-
- if self.server_url is None:
- raise ValueError("Please provide server URL")
-
- if self.model_uid is None:
- raise ValueError("Please provide the model UID")
-
- self._headers: Dict[str, str] = {}
- self._cluster_authed = False
- self._check_cluster_authenticated()
- if api_key is not None and self._cluster_authed:
- self._headers["Authorization"] = f"Bearer {api_key}"
-
- self.client = RESTfulClient(server_url, api_key)
-
- @property
- def _llm_type(self) -> str:
- """Return type of llm."""
- return "xinference"
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- **{"server_url": self.server_url},
- **{"model_uid": self.model_uid},
- **{"model_kwargs": self.model_kwargs},
- }
-
- def _check_cluster_authenticated(self) -> None:
- url = f"{self.server_url}/v1/cluster/auth"
- response = requests.get(url)
- if response.status_code == 404:
- self._cluster_authed = False
- else:
- if response.status_code != 200:
- raise RuntimeError(
- f"Failed to get cluster information, "
- f"detail: {response.json()['detail']}"
- )
- response_data = response.json()
- self._cluster_authed = bool(response_data["auth"])
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call the xinference model and return the output.
-
- Args:
- prompt: The prompt to use for generation.
- stop: Optional list of stop words to use when generating.
- generate_config: Optional dictionary for the configuration used for
- generation.
-
- Returns:
- The generated string by the model.
- """
- if self.client is None:
- raise ValueError("Client is not initialized!")
- model = self.client.get_model(self.model_uid)
-
- generate_config: "LlamaCppGenerateConfig" = kwargs.get("generate_config", {})
-
- generate_config = {**self.model_kwargs, **generate_config}
-
- if stop:
- generate_config["stop"] = stop
-
- if generate_config and generate_config.get("stream"):
- combined_text_output = ""
- for token in self._stream_generate(
- model=model,
- prompt=prompt,
- run_manager=run_manager,
- generate_config=generate_config,
- ):
- combined_text_output += token
- return combined_text_output
-
- else:
- completion = model.generate(prompt=prompt, generate_config=generate_config)
- return completion["choices"][0]["text"]
-
- def _stream_generate(
- self,
- model: Union["RESTfulGenerateModelHandle", "RESTfulChatModelHandle"],
- prompt: str,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- generate_config: Optional["LlamaCppGenerateConfig"] = None,
- ) -> Generator[str, None, None]:
- """
- Args:
- prompt: The prompt to use for generation.
- model: The model used for generation.
- stop: Optional list of stop words to use when generating.
- generate_config: Optional dictionary for the configuration used for
- generation.
-
- Yields:
- A string token.
- """
- streaming_response = model.generate(
- prompt=prompt, generate_config=generate_config
- )
- for chunk in streaming_response:
- if isinstance(chunk, dict):
- choices = chunk.get("choices", [])
- if choices:
- choice = choices[0]
- if isinstance(choice, dict):
- token = choice.get("text", "")
- log_probs = choice.get("logprobs")
- if run_manager:
- run_manager.on_llm_new_token(
- token=token, verbose=self.verbose, log_probs=log_probs
- )
- yield token
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- generate_config = kwargs.get("generate_config", {})
- generate_config = {**self.model_kwargs, **generate_config}
- if stop:
- generate_config["stop"] = stop
- for stream_resp in self._create_generate_stream(prompt, generate_config):
- if stream_resp:
- chunk = self._stream_response_to_generation_chunk(stream_resp)
- if run_manager:
- run_manager.on_llm_new_token(
- chunk.text,
- verbose=self.verbose,
- )
- yield chunk
-
- def _create_generate_stream(
- self, prompt: str, generate_config: Optional[Dict[str, List[str]]] = None
- ) -> Iterator[str]:
- if self.client is None:
- raise ValueError("Client is not initialized!")
- model = self.client.get_model(self.model_uid)
- yield from model.generate(prompt=prompt, generate_config=generate_config)
-
- @staticmethod
- def _stream_response_to_generation_chunk(
- stream_response: str,
- ) -> GenerationChunk:
- """Convert a stream response to a generation chunk."""
- token = ""
- if isinstance(stream_response, dict):
- choices = stream_response.get("choices", [])
- if choices:
- choice = choices[0]
- if isinstance(choice, dict):
- token = choice.get("text", "")
-
- return GenerationChunk(
- text=token,
- generation_info=dict(
- finish_reason=choice.get("finish_reason", None),
- logprobs=choice.get("logprobs", None),
- ),
- )
- else:
- raise TypeError("choice type error!")
- else:
- return GenerationChunk(text=token)
- else:
- raise TypeError("stream_response type error!")
-
- async def _astream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> AsyncIterator[GenerationChunk]:
- generate_config = kwargs.get("generate_config", {})
- generate_config = {**self.model_kwargs, **generate_config}
- if stop:
- generate_config["stop"] = stop
- async for stream_resp in self._acreate_generate_stream(prompt, generate_config):
- if stream_resp:
- chunk = self._stream_response_to_generation_chunk(stream_resp)
- if run_manager:
- await run_manager.on_llm_new_token(
- chunk.text,
- verbose=self.verbose,
- )
- yield chunk
-
- async def _acreate_generate_stream(
- self, prompt: str, generate_config: Optional[Dict[str, List[str]]] = None
- ) -> AsyncIterator[str]:
- request_body: Dict[str, Any] = {"model": self.model_uid, "prompt": prompt}
- if generate_config is not None:
- for key, value in generate_config.items():
- request_body[key] = value
-
- stream = bool(generate_config and generate_config.get("stream"))
- async with aiohttp.ClientSession() as session:
- async with session.post(
- url=f"{self.server_url}/v1/completions",
- json=request_body,
- ) as response:
- if response.status != 200:
- if response.status == 404:
- raise FileNotFoundError(
- "astream call failed with status code 404."
- )
- else:
- optional_detail = response.text
- raise ValueError(
- f"astream call failed with status code {response.status}."
- f" Details: {optional_detail}"
- )
-
- async for line in response.content:
- if not stream:
- yield json.loads(line)
- else:
- json_str = line.decode("utf-8")
- if line.startswith(b"data:"):
- json_str = json_str[len(b"data:") :].strip()
- if not json_str:
- continue
- yield json.loads(json_str)
diff --git a/libs/community/langchain_community/llms/yandex.py b/libs/community/langchain_community/llms/yandex.py
deleted file mode 100644
index 31b09bcdb2..0000000000
--- a/libs/community/langchain_community/llms/yandex.py
+++ /dev/null
@@ -1,345 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Callable, Dict, List, Optional, Sequence
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForLLMRun,
- CallbackManagerForLLMRun,
-)
-from langchain_core.language_models.llms import LLM
-from langchain_core.load.serializable import Serializable
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import SecretStr
-from tenacity import (
- before_sleep_log,
- retry,
- retry_if_exception_type,
- stop_after_attempt,
- wait_exponential,
-)
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-class _BaseYandexGPT(Serializable):
- iam_token: SecretStr = "" # type: ignore[assignment]
- """Yandex Cloud IAM token for service or user account
- with the `ai.languageModels.user` role"""
- api_key: SecretStr = "" # type: ignore[assignment]
- """Yandex Cloud Api Key for service account
- with the `ai.languageModels.user` role"""
- folder_id: str = ""
- """Yandex Cloud folder ID"""
- model_uri: str = ""
- """Model uri to use."""
- model_name: str = "yandexgpt-lite"
- """Model name to use."""
- model_version: str = "latest"
- """Model version to use."""
- temperature: float = 0.6
- """What sampling temperature to use.
- Should be a double number between 0 (inclusive) and 1 (inclusive)."""
- max_tokens: int = 7400
- """Sets the maximum limit on the total number of tokens
- used for both the input prompt and the generated response.
- Must be greater than zero and not exceed 7400 tokens."""
- stop: Optional[List[str]] = None
- """Sequences when completion generation will stop."""
- url: str = "llm.api.cloud.yandex.net:443"
- """The url of the API."""
- max_retries: int = 6
- """Maximum number of retries to make when generating."""
- sleep_interval: float = 1.0
- """Delay between API requests"""
- disable_request_logging: bool = False
- """YandexGPT API logs all request data by default.
- If you provide personal data, confidential information, disable logging."""
- grpc_metadata: Optional[Sequence] = None
-
- @property
- def _llm_type(self) -> str:
- return "yandex_gpt"
-
- @property
- def _identifying_params(self) -> Dict[str, Any]:
- """Get the identifying parameters."""
- return {
- "model_uri": self.model_uri,
- "temperature": self.temperature,
- "max_tokens": self.max_tokens,
- "stop": self.stop,
- "max_retries": self.max_retries,
- }
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that iam token exists in environment."""
-
- iam_token = convert_to_secret_str(
- get_from_dict_or_env(values, "iam_token", "YC_IAM_TOKEN", "")
- )
- values["iam_token"] = iam_token
- api_key = convert_to_secret_str(
- get_from_dict_or_env(values, "api_key", "YC_API_KEY", "")
- )
- values["api_key"] = api_key
- folder_id = get_from_dict_or_env(values, "folder_id", "YC_FOLDER_ID", "")
- values["folder_id"] = folder_id
- if api_key.get_secret_value() == "" and iam_token.get_secret_value() == "":
- raise ValueError("Either 'YC_API_KEY' or 'YC_IAM_TOKEN' must be provided.")
-
- if values["iam_token"]:
- values["grpc_metadata"] = [
- ("authorization", f"Bearer {values['iam_token'].get_secret_value()}")
- ]
- if values["folder_id"]:
- values["grpc_metadata"].append(("x-folder-id", values["folder_id"]))
- else:
- values["grpc_metadata"] = [
- ("authorization", f"Api-Key {values['api_key'].get_secret_value()}"),
- ]
- if values["model_uri"] == "" and values["folder_id"] == "":
- raise ValueError("Either 'model_uri' or 'folder_id' must be provided.")
- if not values["model_uri"]:
- values["model_uri"] = (
- f"gpt://{values['folder_id']}/{values['model_name']}/{values['model_version']}"
- )
- if values["disable_request_logging"]:
- values["grpc_metadata"].append(
- (
- "x-data-logging-enabled",
- "false",
- )
- )
- return values
-
-
-class YandexGPT(_BaseYandexGPT, LLM):
- """Yandex large language models.
-
- To use, you should have the ``yandexcloud`` python package installed.
-
- There are two authentication options for the service account
- with the ``ai.languageModels.user`` role:
- - You can specify the token in a constructor parameter `iam_token`
- or in an environment variable `YC_IAM_TOKEN`.
- - You can specify the key in a constructor parameter `api_key`
- or in an environment variable `YC_API_KEY`.
-
- To use the default model specify the folder ID in a parameter `folder_id`
- or in an environment variable `YC_FOLDER_ID`.
-
- Or specify the model URI in a constructor parameter `model_uri`
-
- Example:
- .. code-block:: python
-
- from langchain_community.llms import YandexGPT
- yandex_gpt = YandexGPT(iam_token="t1.9eu...", folder_id="b1g...")
- """
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call the Yandex GPT model and return the output.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = YandexGPT("Tell me a joke.")
- """
- text = completion_with_retry(self, prompt=prompt)
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
- return text
-
- async def _acall(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Async call the Yandex GPT model and return the output.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
- """
- text = await acompletion_with_retry(self, prompt=prompt)
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
- return text
-
-
-def _make_request(
- self: YandexGPT,
- prompt: str,
-) -> str:
- try:
- import grpc
- from google.protobuf.wrappers_pb2 import DoubleValue, Int64Value
-
- try:
- from yandex.cloud.ai.foundation_models.v1.text_common_pb2 import (
- CompletionOptions,
- Message,
- )
- from yandex.cloud.ai.foundation_models.v1.text_generation.text_generation_service_pb2 import ( # noqa: E501
- CompletionRequest,
- )
- from yandex.cloud.ai.foundation_models.v1.text_generation.text_generation_service_pb2_grpc import ( # noqa: E501
- TextGenerationServiceStub,
- )
- except ModuleNotFoundError:
- from yandex.cloud.ai.foundation_models.v1.foundation_models_pb2 import (
- CompletionOptions,
- Message,
- )
- from yandex.cloud.ai.foundation_models.v1.foundation_models_service_pb2 import ( # noqa: E501
- CompletionRequest,
- )
- from yandex.cloud.ai.foundation_models.v1.foundation_models_service_pb2_grpc import ( # noqa: E501
- TextGenerationServiceStub,
- )
- except ImportError as e:
- raise ImportError(
- "Please install YandexCloud SDK with `pip install yandexcloud` \
- or upgrade it to recent version."
- ) from e
- channel_credentials = grpc.ssl_channel_credentials()
- channel = grpc.secure_channel(self.url, channel_credentials)
- request = CompletionRequest(
- model_uri=self.model_uri,
- completion_options=CompletionOptions(
- temperature=DoubleValue(value=self.temperature),
- max_tokens=Int64Value(value=self.max_tokens),
- ),
- messages=[Message(role="user", text=prompt)],
- )
- stub = TextGenerationServiceStub(channel)
- res = stub.Completion(request, metadata=self.grpc_metadata)
- return list(res)[0].alternatives[0].message.text
-
-
-async def _amake_request(self: YandexGPT, prompt: str) -> str:
- try:
- import asyncio
-
- import grpc
- from google.protobuf.wrappers_pb2 import DoubleValue, Int64Value
-
- try:
- from yandex.cloud.ai.foundation_models.v1.text_common_pb2 import (
- CompletionOptions,
- Message,
- )
- from yandex.cloud.ai.foundation_models.v1.text_generation.text_generation_service_pb2 import ( # noqa: E501
- CompletionRequest,
- CompletionResponse,
- )
- from yandex.cloud.ai.foundation_models.v1.text_generation.text_generation_service_pb2_grpc import ( # noqa: E501
- TextGenerationAsyncServiceStub,
- )
- except ModuleNotFoundError:
- from yandex.cloud.ai.foundation_models.v1.foundation_models_pb2 import (
- CompletionOptions,
- Message,
- )
- from yandex.cloud.ai.foundation_models.v1.foundation_models_service_pb2 import ( # noqa: E501
- CompletionRequest,
- CompletionResponse,
- )
- from yandex.cloud.ai.foundation_models.v1.foundation_models_service_pb2_grpc import ( # noqa: E501
- TextGenerationAsyncServiceStub,
- )
- from yandex.cloud.operation.operation_service_pb2 import GetOperationRequest
- from yandex.cloud.operation.operation_service_pb2_grpc import (
- OperationServiceStub,
- )
- except ImportError as e:
- raise ImportError(
- "Please install YandexCloud SDK with `pip install yandexcloud` \
- or upgrade it to recent version."
- ) from e
- operation_api_url = "operation.api.cloud.yandex.net:443"
- channel_credentials = grpc.ssl_channel_credentials()
- async with grpc.aio.secure_channel(self.url, channel_credentials) as channel:
- request = CompletionRequest(
- model_uri=self.model_uri,
- completion_options=CompletionOptions(
- temperature=DoubleValue(value=self.temperature),
- max_tokens=Int64Value(value=self.max_tokens),
- ),
- messages=[Message(role="user", text=prompt)],
- )
- stub = TextGenerationAsyncServiceStub(channel)
- operation = await stub.Completion(request, metadata=self.grpc_metadata)
- async with grpc.aio.secure_channel(
- operation_api_url, channel_credentials
- ) as operation_channel:
- operation_stub = OperationServiceStub(operation_channel)
- while not operation.done:
- await asyncio.sleep(1)
- operation_request = GetOperationRequest(operation_id=operation.id)
- operation = await operation_stub.Get(
- operation_request,
- metadata=self.grpc_metadata,
- )
-
- completion_response = CompletionResponse()
- operation.response.Unpack(completion_response)
- return completion_response.alternatives[0].message.text
-
-
-def _create_retry_decorator(llm: YandexGPT) -> Callable[[Any], Any]:
- from grpc import RpcError
-
- min_seconds = llm.sleep_interval
- max_seconds = 60
- return retry(
- reraise=True,
- stop=stop_after_attempt(llm.max_retries),
- wait=wait_exponential(multiplier=1, min=min_seconds, max=max_seconds),
- retry=(retry_if_exception_type((RpcError))),
- before_sleep=before_sleep_log(logger, logging.WARNING),
- )
-
-
-def completion_with_retry(llm: YandexGPT, **kwargs: Any) -> Any:
- """Use tenacity to retry the completion call."""
- retry_decorator = _create_retry_decorator(llm)
-
- @retry_decorator
- def _completion_with_retry(**_kwargs: Any) -> Any:
- return _make_request(llm, **_kwargs)
-
- return _completion_with_retry(**kwargs)
-
-
-async def acompletion_with_retry(llm: YandexGPT, **kwargs: Any) -> Any:
- """Use tenacity to retry the async completion call."""
- retry_decorator = _create_retry_decorator(llm)
-
- @retry_decorator
- async def _completion_with_retry(**_kwargs: Any) -> Any:
- return await _amake_request(llm, **_kwargs)
-
- return await _completion_with_retry(**kwargs)
diff --git a/libs/community/langchain_community/llms/yi.py b/libs/community/langchain_community/llms/yi.py
deleted file mode 100644
index 6f6dc96380..0000000000
--- a/libs/community/langchain_community/llms/yi.py
+++ /dev/null
@@ -1,104 +0,0 @@
-from __future__ import annotations
-
-import json
-import logging
-from typing import Any, Dict, List, Literal, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env
-from pydantic import Field, SecretStr
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-class YiLLM(LLM):
- """Yi large language models."""
-
- model: str = "yi-large"
- temperature: float = 0.3
- top_p: float = 0.95
- timeout: int = 60
- model_kwargs: Dict[str, Any] = Field(default_factory=dict)
-
- yi_api_key: Optional[SecretStr] = None
- region: Literal["auto", "domestic", "international"] = "auto"
- yi_api_url_domestic: str = "https://api.lingyiwanwu.com/v1/chat/completions"
- yi_api_url_international: str = "https://api.01.ai/v1/chat/completions"
-
- def __init__(self, **kwargs: Any):
- kwargs["yi_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(kwargs, "yi_api_key", "YI_API_KEY")
- )
- super().__init__(**kwargs)
-
- @property
- def _default_params(self) -> Dict[str, Any]:
- return {
- "model": self.model,
- "temperature": self.temperature,
- "top_p": self.top_p,
- **self.model_kwargs,
- }
-
- def _post(self, request: Any) -> Any:
- headers = {
- "Content-Type": "application/json",
- "Authorization": f"Bearer {self.yi_api_key.get_secret_value()}", # type: ignore[union-attr]
- }
-
- urls = []
- if self.region == "domestic":
- urls = [self.yi_api_url_domestic]
- elif self.region == "international":
- urls = [self.yi_api_url_international]
- else: # auto
- urls = [self.yi_api_url_domestic, self.yi_api_url_international]
-
- for url in urls:
- try:
- response = requests.post(
- url,
- headers=headers,
- json=request,
- timeout=self.timeout,
- )
-
- if response.status_code == 200:
- parsed_json = json.loads(response.text)
- return parsed_json["choices"][0]["message"]["content"]
- elif (
- response.status_code != 403
- ): # If not a permission error, raise immediately
- response.raise_for_status()
- except requests.RequestException as e:
- if url == urls[-1]: # If this is the last URL to try
- raise ValueError(f"An error has occurred: {e}")
- else:
- logger.warning(f"Failed to connect to {url}, trying next URL")
- continue
-
- raise ValueError("Failed to connect to all available URLs")
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- request = self._default_params
- request["messages"] = [{"role": "user", "content": prompt}]
- request.update(kwargs)
- text = self._post(request)
- if stop is not None:
- text = enforce_stop_tokens(text, stop)
- return text
-
- @property
- def _llm_type(self) -> str:
- """Return type of chat_model."""
- return "yi-llm"
diff --git a/libs/community/langchain_community/llms/you.py b/libs/community/langchain_community/llms/you.py
deleted file mode 100644
index 20ba6ca451..0000000000
--- a/libs/community/langchain_community/llms/you.py
+++ /dev/null
@@ -1,140 +0,0 @@
-import os
-from typing import Any, Dict, Generator, Iterator, List, Literal, Optional
-
-import requests
-from langchain_core.callbacks.manager import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from langchain_core.outputs import GenerationChunk
-from pydantic import Field
-
-SMART_ENDPOINT = "https://chat-api.you.com/smart"
-RESEARCH_ENDPOINT = "https://chat-api.you.com/research"
-
-
-def _request(base_url: str, api_key: str, **kwargs: Any) -> Dict[str, Any]:
- """
- NOTE: This function can be replaced by a OpenAPI-generated Python SDK in the future,
- for better input/output typing support.
- """
- headers = {"x-api-key": api_key}
- response = requests.post(base_url, headers=headers, json=kwargs)
- response.raise_for_status()
- return response.json()
-
-
-def _request_stream(
- base_url: str, api_key: str, **kwargs: Any
-) -> Generator[str, None, None]:
- headers = {"x-api-key": api_key}
- params = dict(**kwargs, stream=True)
- response = requests.post(base_url, headers=headers, stream=True, json=params)
- response.raise_for_status()
-
- # Explicitly coercing the response to a generator to satisfy mypy
- event_source = (bytestring for bytestring in response)
-
- try:
- import sseclient
-
- client = sseclient.SSEClient(event_source)
- except ImportError:
- raise ImportError(
- (
- "Could not import `sseclient`. "
- "Please install it with `pip install sseclient-py`."
- )
- )
-
- for event in client.events():
- if event.event in ("search_results", "done"):
- pass
- elif event.event == "token":
- yield event.data
- elif event.event == "error":
- raise ValueError(f"Error in response: {event.data}")
- else:
- raise NotImplementedError(f"Unknown event type {event.event}")
-
-
-class You(LLM):
- """Wrapper around You.com's conversational Smart and Research APIs.
-
- Each API endpoint is designed to generate conversational
- responses to a variety of query types, including inline citations
- and web results when relevant.
-
- Smart Endpoint:
- - Quick, reliable answers for a variety of questions
- - Cites the entire web page URL
-
- Research Endpoint:
- - In-depth answers with extensive citations for a variety of questions
- - Cites the specific web page snippet relevant to the claim
-
- To connect to the You.com api requires an API key which
- you can get at https://api.you.com.
-
- For more information, check out the documentations at
- https://documentation.you.com/api-reference/.
-
- Args:
- endpoint: You.com conversational endpoints. Choose from "smart" or "research"
- ydc_api_key: You.com API key, if `YDC_API_KEY` is not set in the environment
- """
-
- endpoint: Literal["smart", "research"] = Field(
- "smart",
- description=(
- 'You.com conversational endpoints. Choose from "smart" or "research"'
- ),
- )
- ydc_api_key: Optional[str] = Field(
- None,
- description="You.com API key, if `YDC_API_KEY` is not set in the envrioment",
- )
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- if stop:
- raise NotImplementedError(
- "Stop words are not implemented for You.com endpoints."
- )
- params = {"query": prompt}
- response = _request(self._request_endpoint, api_key=self._api_key, **params)
- return response["answer"]
-
- def _stream(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> Iterator[GenerationChunk]:
- if stop:
- raise NotImplementedError(
- "Stop words are not implemented for You.com endpoints."
- )
- params = {"query": prompt}
- for token in _request_stream(
- self._request_endpoint, api_key=self._api_key, **params
- ):
- yield GenerationChunk(text=token)
-
- @property
- def _request_endpoint(self) -> str:
- if self.endpoint == "smart":
- return SMART_ENDPOINT
- return RESEARCH_ENDPOINT
-
- @property
- def _api_key(self) -> str:
- return self.ydc_api_key or os.environ["YDC_API_KEY"]
-
- @property
- def _llm_type(self) -> str:
- return "you.com"
diff --git a/libs/community/langchain_community/llms/yuan2.py b/libs/community/langchain_community/llms/yuan2.py
deleted file mode 100644
index 1087d0cb00..0000000000
--- a/libs/community/langchain_community/llms/yuan2.py
+++ /dev/null
@@ -1,205 +0,0 @@
-import json
-import logging
-from typing import Any, Dict, List, Mapping, Optional, Set
-
-import requests
-from langchain_core.callbacks import CallbackManagerForLLMRun
-from langchain_core.language_models.llms import LLM
-from pydantic import Field
-
-from langchain_community.llms.utils import enforce_stop_tokens
-
-logger = logging.getLogger(__name__)
-
-
-class Yuan2(LLM):
- """Yuan2.0 language models.
-
- Example:
- .. code-block:: python
-
- yuan_llm = Yuan2(
- infer_api="http://127.0.0.1:8000/yuan",
- max_tokens=1024,
- temp=1.0,
- top_p=0.9,
- top_k=40,
- )
- print(yuan_llm)
- print(yuan_llm.invoke("你是谁?"))
- """
-
- infer_api: str = "http://127.0.0.1:8000/yuan"
- """Yuan2.0 inference api"""
-
- max_tokens: int = Field(1024, alias="max_token")
- """Token context window."""
-
- temp: Optional[float] = 0.7
- """The temperature to use for sampling."""
-
- top_p: Optional[float] = 0.9
- """The top-p value to use for sampling."""
-
- top_k: Optional[int] = 0
- """The top-k value to use for sampling."""
-
- do_sample: bool = False
- """The do_sample is a Boolean value that determines whether
- to use the sampling method during text generation.
- """
-
- echo: Optional[bool] = False
- """Whether to echo the prompt."""
-
- stop: Optional[List[str]] = []
- """A list of strings to stop generation when encountered."""
-
- repeat_last_n: Optional[int] = 64
- "Last n tokens to penalize"
-
- repeat_penalty: Optional[float] = 1.18
- """The penalty to apply to repeated tokens."""
-
- streaming: bool = False
- """Whether to stream the results or not."""
-
- history: List[str] = []
- """History of the conversation"""
-
- use_history: bool = False
- """Whether to use history or not"""
-
- def __init__(self, **kwargs: Any) -> None:
- """Initialize the Yuan2 class."""
- super().__init__(**kwargs)
-
- if (self.top_p or 0) > 0 and (self.top_k or 0) > 0:
- logger.warning(
- "top_p and top_k cannot be set simultaneously. "
- "set top_k to 0 instead..."
- )
- self.top_k = 0
-
- @property
- def _llm_type(self) -> str:
- return "Yuan2.0"
-
- @staticmethod
- def _model_param_names() -> Set[str]:
- return {
- "max_tokens",
- "temp",
- "top_k",
- "top_p",
- "do_sample",
- }
-
- def _default_params(self) -> Dict[str, Any]:
- return {
- "do_sample": self.do_sample,
- "infer_api": self.infer_api,
- "max_tokens": self.max_tokens,
- "repeat_penalty": self.repeat_penalty,
- "temp": self.temp,
- "top_k": self.top_k,
- "top_p": self.top_p,
- "use_history": self.use_history,
- }
-
- @property
- def _identifying_params(self) -> Mapping[str, Any]:
- """Get the identifying parameters."""
- return {
- "model": self._llm_type,
- **self._default_params(),
- **{
- k: v for k, v in self.__dict__.items() if k in self._model_param_names()
- },
- }
-
- def _call(
- self,
- prompt: str,
- stop: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForLLMRun] = None,
- **kwargs: Any,
- ) -> str:
- """Call out to a Yuan2.0 LLM inference endpoint.
-
- Args:
- prompt: The prompt to pass into the model.
- stop: Optional list of stop words to use when generating.
-
- Returns:
- The string generated by the model.
-
- Example:
- .. code-block:: python
-
- response = yuan_llm.invoke("你能做什么?")
- """
-
- if self.use_history:
- self.history.append(prompt)
- input = "".join(self.history)
- else:
- input = prompt
-
- headers = {"Content-Type": "application/json"}
-
- data = json.dumps(
- {
- "ques_list": [{"id": "000", "ques": input}],
- "tokens_to_generate": self.max_tokens,
- "temperature": self.temp,
- "top_p": self.top_p,
- "top_k": self.top_k,
- "do_sample": self.do_sample,
- }
- )
-
- logger.debug("Yuan2.0 prompt:", input)
-
- # call api
- try:
- response = requests.put(self.infer_api, headers=headers, data=data)
- except requests.exceptions.RequestException as e:
- raise ValueError(f"Error raised by inference api: {e}")
-
- logger.debug(f"Yuan2.0 response: {response}")
-
- if response.status_code != 200:
- raise ValueError(f"Failed with response: {response}")
- try:
- resp = response.json()
-
- if resp["errCode"] != "0":
- raise ValueError(
- f"Failed with error code [{resp['errCode']}], "
- f"error message: [{resp['exceptionMsg']}]"
- )
-
- if "resData" in resp:
- if len(resp["resData"]["output"]) >= 0:
- generate_text = resp["resData"]["output"][0]["ans"]
- else:
- raise ValueError("No output found in response.")
- else:
- raise ValueError("No resData found in response.")
-
- except requests.exceptions.JSONDecodeError as e:
- raise ValueError(
- f"Error raised during decoding response from inference api: {e}."
- f"\nResponse: {response.text}"
- )
-
- if stop is not None:
- generate_text = enforce_stop_tokens(generate_text, stop)
-
- # support multi-turn chat
- if self.use_history:
- self.history.append(generate_text)
-
- logger.debug(f"history: {self.history}")
- return generate_text
diff --git a/libs/community/langchain_community/memory/__init__.py b/libs/community/langchain_community/memory/__init__.py
deleted file mode 100644
index e69de29bb2..0000000000
diff --git a/libs/community/langchain_community/memory/kg.py b/libs/community/langchain_community/memory/kg.py
deleted file mode 100644
index f60e4f5b75..0000000000
--- a/libs/community/langchain_community/memory/kg.py
+++ /dev/null
@@ -1,141 +0,0 @@
-from typing import Any, Dict, List, Type, Union
-
-from langchain_core.language_models import BaseLanguageModel
-from langchain_core.messages import BaseMessage, SystemMessage, get_buffer_string
-from langchain_core.prompts import BasePromptTemplate
-from pydantic import Field
-
-from langchain_community.graphs import NetworkxEntityGraph
-from langchain_community.graphs.networkx_graph import (
- KnowledgeTriple,
- get_entities,
- parse_triples,
-)
-
-try:
- from langchain.chains.llm import LLMChain
- from langchain.memory.chat_memory import BaseChatMemory
- from langchain.memory.prompt import (
- ENTITY_EXTRACTION_PROMPT,
- KNOWLEDGE_TRIPLE_EXTRACTION_PROMPT,
- )
- from langchain.memory.utils import get_prompt_input_key
-
- class ConversationKGMemory(BaseChatMemory):
- """Knowledge graph conversation memory.
-
- Integrates with external knowledge graph to store and retrieve
- information about knowledge triples in the conversation.
- """
-
- k: int = 2
- human_prefix: str = "Human"
- ai_prefix: str = "AI"
- kg: NetworkxEntityGraph = Field(default_factory=NetworkxEntityGraph)
- knowledge_extraction_prompt: BasePromptTemplate = (
- KNOWLEDGE_TRIPLE_EXTRACTION_PROMPT
- )
- entity_extraction_prompt: BasePromptTemplate = ENTITY_EXTRACTION_PROMPT
- llm: BaseLanguageModel
- summary_message_cls: Type[BaseMessage] = SystemMessage
- """Number of previous utterances to include in the context."""
- memory_key: str = "history" #: :meta private:
-
- def load_memory_variables(self, inputs: Dict[str, Any]) -> Dict[str, Any]:
- """Return history buffer."""
- entities = self._get_current_entities(inputs)
-
- summary_strings = []
- for entity in entities:
- knowledge = self.kg.get_entity_knowledge(entity)
- if knowledge:
- summary = f"On {entity}: {'. '.join(knowledge)}."
- summary_strings.append(summary)
- context: Union[str, List]
- if not summary_strings:
- context = [] if self.return_messages else ""
- elif self.return_messages:
- context = [
- self.summary_message_cls(content=text) for text in summary_strings
- ]
- else:
- context = "\n".join(summary_strings)
-
- return {self.memory_key: context}
-
- @property
- def memory_variables(self) -> List[str]:
- """Will always return list of memory variables.
-
- :meta private:
- """
- return [self.memory_key]
-
- def _get_prompt_input_key(self, inputs: Dict[str, Any]) -> str:
- """Get the input key for the prompt."""
- if self.input_key is None:
- return get_prompt_input_key(inputs, self.memory_variables)
- return self.input_key
-
- def _get_prompt_output_key(self, outputs: Dict[str, Any]) -> str:
- """Get the output key for the prompt."""
- if self.output_key is None:
- if len(outputs) != 1:
- raise ValueError(f"One output key expected, got {outputs.keys()}")
- return list(outputs.keys())[0]
- return self.output_key
-
- def get_current_entities(self, input_string: str) -> List[str]:
- chain = LLMChain(llm=self.llm, prompt=self.entity_extraction_prompt)
- buffer_string = get_buffer_string(
- self.chat_memory.messages[-self.k * 2 :],
- human_prefix=self.human_prefix,
- ai_prefix=self.ai_prefix,
- )
- output = chain.predict(
- history=buffer_string,
- input=input_string,
- )
- return get_entities(output)
-
- def _get_current_entities(self, inputs: Dict[str, Any]) -> List[str]:
- """Get the current entities in the conversation."""
- prompt_input_key = self._get_prompt_input_key(inputs)
- return self.get_current_entities(inputs[prompt_input_key])
-
- def get_knowledge_triplets(self, input_string: str) -> List[KnowledgeTriple]:
- chain = LLMChain(llm=self.llm, prompt=self.knowledge_extraction_prompt)
- buffer_string = get_buffer_string(
- self.chat_memory.messages[-self.k * 2 :],
- human_prefix=self.human_prefix,
- ai_prefix=self.ai_prefix,
- )
- output = chain.predict(
- history=buffer_string,
- input=input_string,
- verbose=True,
- )
- knowledge = parse_triples(output)
- return knowledge
-
- def _get_and_update_kg(self, inputs: Dict[str, Any]) -> None:
- """Get and update knowledge graph from the conversation history."""
- prompt_input_key = self._get_prompt_input_key(inputs)
- knowledge = self.get_knowledge_triplets(inputs[prompt_input_key])
- for triple in knowledge:
- self.kg.add_triple(triple)
-
- def save_context(self, inputs: Dict[str, Any], outputs: Dict[str, str]) -> None:
- """Save context from this conversation to buffer."""
- super().save_context(inputs, outputs)
- self._get_and_update_kg(inputs)
-
- def clear(self) -> None:
- """Clear memory contents."""
- super().clear()
- self.kg.clear()
-
-except ImportError:
- # Placeholder object
- class ConversationKGMemory: # type: ignore[no-redef]
- pass
diff --git a/libs/community/langchain_community/memory/motorhead_memory.py b/libs/community/langchain_community/memory/motorhead_memory.py
deleted file mode 100644
index 3b41be424e..0000000000
--- a/libs/community/langchain_community/memory/motorhead_memory.py
+++ /dev/null
@@ -1,101 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-import requests
-from langchain_core.messages import get_buffer_string
-
-try:
- # Temporarily tuck import in a conditional import until
- # community pkg becomes dependent on langchain core
- from langchain.memory.chat_memory import BaseChatMemory
-
- MANAGED_URL = "https://api.getmetal.io/v1/motorhead"
-
- class MotorheadMemory(BaseChatMemory):
- """Chat message memory backed by Motorhead service."""
-
- url: str = MANAGED_URL
- timeout: int = 3000
- memory_key: str = "history"
- session_id: str
- context: Optional[str] = None
-
- # Managed Params
- api_key: Optional[str] = None
- client_id: Optional[str] = None
-
- def __get_headers(self) -> Dict[str, str]:
- is_managed = self.url == MANAGED_URL
-
- headers = {
- "Content-Type": "application/json",
- }
-
- if is_managed and not (self.api_key and self.client_id):
- raise ValueError(
- """
- You must provide an API key or a client ID to use the managed
- version of Motorhead. Visit https://getmetal.io
- for more information.
- """
- )
-
- if is_managed and self.api_key and self.client_id:
- headers["x-metal-api-key"] = self.api_key
- headers["x-metal-client-id"] = self.client_id
-
- return headers
-
- async def init(self) -> None:
- res = requests.get(
- f"{self.url}/sessions/{self.session_id}/memory",
- timeout=self.timeout,
- headers=self.__get_headers(),
- )
- res_data = res.json()
- res_data = res_data.get("data", res_data) # Handle Managed Version
-
- messages = res_data.get("messages", [])
- context = res_data.get("context", "NONE")
-
- for message in reversed(messages):
- if message["role"] == "AI":
- self.chat_memory.add_ai_message(message["content"])
- else:
- self.chat_memory.add_user_message(message["content"])
-
- if context and context != "NONE":
- self.context = context
-
- def load_memory_variables(self, values: Dict[str, Any]) -> Dict[str, Any]:
- if self.return_messages:
- return {self.memory_key: self.chat_memory.messages}
- else:
- return {self.memory_key: get_buffer_string(self.chat_memory.messages)}
-
- @property
- def memory_variables(self) -> List[str]:
- return [self.memory_key]
-
- def save_context(self, inputs: Dict[str, Any], outputs: Dict[str, str]) -> None:
- input_str, output_str = self._get_input_output(inputs, outputs)
- requests.post(
- f"{self.url}/sessions/{self.session_id}/memory",
- timeout=self.timeout,
- json={
- "messages": [
- {"role": "Human", "content": f"{input_str}"},
- {"role": "AI", "content": f"{output_str}"},
- ]
- },
- headers=self.__get_headers(),
- )
- super().save_context(inputs, outputs)
-
- def delete_session(self) -> None:
- """Delete a session"""
- requests.delete(f"{self.url}/sessions/{self.session_id}/memory")
-
-except ImportError:
- # Placeholder object
- class MotorheadMemory: # type: ignore[no-redef]
- pass
diff --git a/libs/community/langchain_community/memory/zep_cloud_memory.py b/libs/community/langchain_community/memory/zep_cloud_memory.py
deleted file mode 100644
index eacab5eb27..0000000000
--- a/libs/community/langchain_community/memory/zep_cloud_memory.py
+++ /dev/null
@@ -1,125 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Dict, Optional
-
-from langchain_community.chat_message_histories import ZepCloudChatMessageHistory
-
-try:
- from langchain.memory import ConversationBufferMemory
- from zep_cloud import MemoryGetRequestMemoryType
-
- class ZepCloudMemory(ConversationBufferMemory):
- """Persist your chain history to the Zep MemoryStore.
-
- Documentation: https://help.getzep.com
-
- Example:
- .. code-block:: python
-
- memory = ZepCloudMemory(
- session_id=session_id, # Identifies your user or a user's session
- api_key=, # Your Zep Project API key
- memory_key="history", # Ensure this matches the key used in
- # chain's prompt template
- return_messages=True, # Does your prompt template expect a string
- # or a list of Messages?
- )
- chain = LLMChain(memory=memory,...) # Configure your chain to use the ZepMemory
- instance
-
-
- Note:
- To persist metadata alongside your chat history, your will need to create a
- custom Chain class that overrides the `prep_outputs` method to include the metadata
- in the call to `self.memory.save_context`.
-
-
- Zep - Recall, understand, and extract data from chat histories. Power personalized AI experiences.
- =========
- Zep is a long-term memory service for AI Assistant apps. With Zep, you can provide AI assistants with the ability to recall past conversations,
- no matter how distant, while also reducing hallucinations, latency, and cost.
-
- For more information on the zep-python package, see:
- https://github.com/getzep/zep-python
-
- """ # noqa: E501
-
- chat_memory: ZepCloudChatMessageHistory
-
- def __init__(
- self,
- session_id: str,
- api_key: str,
- memory_type: Optional[MemoryGetRequestMemoryType] = None,
- lastn: Optional[int] = None,
- output_key: Optional[str] = None,
- input_key: Optional[str] = None,
- return_messages: bool = False,
- human_prefix: str = "Human",
- ai_prefix: str = "AI",
- memory_key: str = "history",
- ):
- """Initialize ZepMemory.
-
- Args:
- session_id (str): Identifies your user or a user's session
- api_key (str): Your Zep Project key.
- memory_type (Optional[MemoryGetRequestMemoryType], optional): Zep Memory Type, defaults to perpetual
- lastn (Optional[int], optional): Number of messages to retrieve. Will add the last summary generated prior to the nth oldest message. Defaults to 6
- output_key (Optional[str], optional): The key to use for the output message.
- Defaults to None.
- input_key (Optional[str], optional): The key to use for the input message.
- Defaults to None.
- return_messages (bool, optional): Does your prompt template expect a string
- or a list of Messages? Defaults to False
- i.e. return a string.
- human_prefix (str, optional): The prefix to use for human messages.
- Defaults to "Human".
- ai_prefix (str, optional): The prefix to use for AI messages.
- Defaults to "AI".
- memory_key (str, optional): The key to use for the memory.
- Defaults to "history".
- Ensure that this matches the key used in
- chain's prompt template.
- """ # noqa: E501
- chat_message_history = ZepCloudChatMessageHistory(
- session_id=session_id,
- memory_type=memory_type,
- lastn=lastn,
- api_key=api_key,
- )
- super().__init__(
- chat_memory=chat_message_history,
- output_key=output_key,
- input_key=input_key,
- return_messages=return_messages,
- human_prefix=human_prefix,
- ai_prefix=ai_prefix,
- memory_key=memory_key,
- )
-
- def save_context(
- self,
- inputs: Dict[str, Any],
- outputs: Dict[str, str],
- metadata: Optional[Dict[str, Any]] = None,
- ) -> None:
- """Save context from this conversation to buffer.
-
- Args:
- inputs (Dict[str, Any]): The inputs to the chain.
- outputs (Dict[str, str]): The outputs from the chain.
- metadata (Optional[Dict[str, Any]], optional): Any metadata to save with
- the context. Defaults to None
-
- Returns:
- None
- """
- input_str, output_str = self._get_input_output(inputs, outputs)
- self.chat_memory.add_user_message(input_str, metadata=metadata)
- self.chat_memory.add_ai_message(output_str, metadata=metadata)
-
-except ImportError:
- # Placeholder object
- class ZepCloudMemory: # type: ignore[no-redef]
- pass
diff --git a/libs/community/langchain_community/memory/zep_memory.py b/libs/community/langchain_community/memory/zep_memory.py
deleted file mode 100644
index eda8042ecb..0000000000
--- a/libs/community/langchain_community/memory/zep_memory.py
+++ /dev/null
@@ -1,130 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Dict, Optional
-
-from langchain_community.chat_message_histories import ZepChatMessageHistory
-
-try:
- from langchain.memory import ConversationBufferMemory
-
- class ZepMemory(ConversationBufferMemory):
- """Persist your chain history to the Zep MemoryStore.
-
- The number of messages returned by Zep and when the Zep server summarizes chat
- histories is configurable. See the Zep documentation for more details.
-
- Documentation: https://docs.getzep.com
-
- Example:
- .. code-block:: python
-
- memory = ZepMemory(
- session_id=session_id, # Identifies your user or a user's session
- url=ZEP_API_URL, # Your Zep server's URL
- api_key=, # Optional
- memory_key="history", # Ensure this matches the key used in
- # chain's prompt template
- return_messages=True, # Does your prompt template expect a string
- # or a list of Messages?
- )
- chain = LLMChain(memory=memory,...) # Configure your chain to use the ZepMemory
- instance
-
-
- Note:
- To persist metadata alongside your chat history, your will need to create a
- custom Chain class that overrides the `prep_outputs` method to include the metadata
- in the call to `self.memory.save_context`.
-
-
- Zep - Fast, scalable building blocks for LLM Apps
- =========
- Zep is an open source platform for productionizing LLM apps. Go from a prototype
- built in LangChain or LlamaIndex, or a custom app, to production in minutes without
- rewriting code.
-
- For server installation instructions and more, see:
- https://docs.getzep.com/deployment/quickstart/
-
- For more information on the zep-python package, see:
- https://github.com/getzep/zep-python
-
- """ # noqa: E501
-
- chat_memory: ZepChatMessageHistory
-
- def __init__(
- self,
- session_id: str,
- url: str = "http://localhost:8000",
- api_key: Optional[str] = None,
- output_key: Optional[str] = None,
- input_key: Optional[str] = None,
- return_messages: bool = False,
- human_prefix: str = "Human",
- ai_prefix: str = "AI",
- memory_key: str = "history",
- ):
- """Initialize ZepMemory.
-
- Args:
- session_id (str): Identifies your user or a user's session
- url (str, optional): Your Zep server's URL. Defaults to
- "http://localhost:8000".
- api_key (Optional[str], optional): Your Zep API key. Defaults to None.
- output_key (Optional[str], optional): The key to use for the output message.
- Defaults to None.
- input_key (Optional[str], optional): The key to use for the input message.
- Defaults to None.
- return_messages (bool, optional): Does your prompt template expect a string
- or a list of Messages? Defaults to False
- i.e. return a string.
- human_prefix (str, optional): The prefix to use for human messages.
- Defaults to "Human".
- ai_prefix (str, optional): The prefix to use for AI messages.
- Defaults to "AI".
- memory_key (str, optional): The key to use for the memory.
- Defaults to "history".
- Ensure that this matches the key used in
- chain's prompt template.
- """ # noqa: E501
- chat_message_history = ZepChatMessageHistory(
- session_id=session_id,
- url=url,
- api_key=api_key,
- )
- super().__init__(
- chat_memory=chat_message_history,
- output_key=output_key,
- input_key=input_key,
- return_messages=return_messages,
- human_prefix=human_prefix,
- ai_prefix=ai_prefix,
- memory_key=memory_key,
- )
-
- def save_context(
- self,
- inputs: Dict[str, Any],
- outputs: Dict[str, str],
- metadata: Optional[Dict[str, Any]] = None,
- ) -> None:
- """Save context from this conversation to buffer.
-
- Args:
- inputs (Dict[str, Any]): The inputs to the chain.
- outputs (Dict[str, str]): The outputs from the chain.
- metadata (Optional[Dict[str, Any]], optional): Any metadata to save with
- the context. Defaults to None
-
- Returns:
- None
- """
- input_str, output_str = self._get_input_output(inputs, outputs)
- self.chat_memory.add_user_message(input_str, metadata=metadata)
- self.chat_memory.add_ai_message(output_str, metadata=metadata)
-
-except ImportError:
- # Placeholder object
- class ZepMemory: # type: ignore[no-redef]
- pass
diff --git a/libs/community/langchain_community/output_parsers/__init__.py b/libs/community/langchain_community/output_parsers/__init__.py
deleted file mode 100644
index 62740af544..0000000000
--- a/libs/community/langchain_community/output_parsers/__init__.py
+++ /dev/null
@@ -1,14 +0,0 @@
-"""**OutputParser** classes parse the output of an LLM call.
-
-**Class hierarchy:**
-
-.. code-block::
-
- BaseLLMOutputParser --> BaseOutputParser --> OutputParser # GuardrailsOutputParser
-
-**Main helpers:**
-
-.. code-block::
-
- Serializable, Generation, PromptValue
-""" # noqa: E501
diff --git a/libs/community/langchain_community/output_parsers/ernie_functions.py b/libs/community/langchain_community/output_parsers/ernie_functions.py
deleted file mode 100644
index 4b23e9e591..0000000000
--- a/libs/community/langchain_community/output_parsers/ernie_functions.py
+++ /dev/null
@@ -1,180 +0,0 @@
-import copy
-import json
-from typing import Any, Dict, List, Optional, Type, Union
-
-import jsonpatch
-from langchain_core.exceptions import OutputParserException
-from langchain_core.output_parsers import (
- BaseCumulativeTransformOutputParser,
- BaseGenerationOutputParser,
-)
-from langchain_core.output_parsers.json import parse_partial_json
-from langchain_core.outputs.chat_generation import (
- ChatGeneration,
- Generation,
-)
-from pydantic import BaseModel, model_validator
-
-
-class OutputFunctionsParser(BaseGenerationOutputParser[Any]):
- """Parse an output that is one of sets of values."""
-
- args_only: bool = True
- """Whether to only return the arguments to the function call."""
-
- def parse_result(self, result: List[Generation], *, partial: bool = False) -> Any:
- generation = result[0]
- if not isinstance(generation, ChatGeneration):
- raise OutputParserException(
- "This output parser can only be used with a chat generation."
- )
- message = generation.message
- try:
- func_call = copy.deepcopy(message.additional_kwargs["function_call"])
- except KeyError as exc:
- raise OutputParserException(f"Could not parse function call: {exc}")
-
- if self.args_only:
- return func_call["arguments"]
- return func_call
-
-
-class JsonOutputFunctionsParser(BaseCumulativeTransformOutputParser[Any]):
- """Parse an output as the Json object."""
-
- strict: bool = False
- """Whether to allow non-JSON-compliant strings.
-
- See: https://docs.python.org/3/library/json.html#encoders-and-decoders
-
- Useful when the parsed output may include unicode characters or new lines.
- """
-
- args_only: bool = True
- """Whether to only return the arguments to the function call."""
-
- @property
- def _type(self) -> str:
- return "json_functions"
-
- def _diff(self, prev: Optional[Any], next: Any) -> Any:
- return jsonpatch.make_patch(prev, next).patch
-
- def parse_result(self, result: List[Generation], *, partial: bool = False) -> Any:
- if len(result) != 1:
- raise OutputParserException(
- f"Expected exactly one result, but got {len(result)}"
- )
- generation = result[0]
- if not isinstance(generation, ChatGeneration):
- raise OutputParserException(
- "This output parser can only be used with a chat generation."
- )
- message = generation.message
- if "function_call" not in message.additional_kwargs:
- return None
- try:
- function_call = message.additional_kwargs["function_call"]
- except KeyError as exc:
- if partial:
- return None
- else:
- raise OutputParserException(f"Could not parse function call: {exc}")
- try:
- if partial:
- if self.args_only:
- return parse_partial_json(
- function_call["arguments"], strict=self.strict
- )
- else:
- return {
- **function_call,
- "arguments": parse_partial_json(
- function_call["arguments"], strict=self.strict
- ),
- }
- else:
- if self.args_only:
- try:
- return json.loads(
- function_call["arguments"], strict=self.strict
- )
- except (json.JSONDecodeError, TypeError) as exc:
- raise OutputParserException(
- f"Could not parse function call data: {exc}"
- )
- else:
- try:
- return {
- **function_call,
- "arguments": json.loads(
- function_call["arguments"], strict=self.strict
- ),
- }
- except (json.JSONDecodeError, TypeError) as exc:
- raise OutputParserException(
- f"Could not parse function call data: {exc}"
- )
- except KeyError:
- return None
-
- # This method would be called by the default implementation of `parse_result`
- # but we're overriding that method so it's not needed.
- def parse(self, text: str) -> Any:
- raise NotImplementedError()
-
-
-class JsonKeyOutputFunctionsParser(JsonOutputFunctionsParser):
- """Parse an output as the element of the Json object."""
-
- key_name: str
- """The name of the key to return."""
-
- def parse_result(self, result: List[Generation], *, partial: bool = False) -> Any:
- res = super().parse_result(result, partial=partial)
- if partial and res is None:
- return None
- return res.get(self.key_name) if partial else res[self.key_name]
-
-
-class PydanticOutputFunctionsParser(OutputFunctionsParser):
- """Parse an output as a pydantic object."""
-
- pydantic_schema: Union[Type[BaseModel], Dict[str, Type[BaseModel]]]
- """The pydantic schema to parse the output with."""
-
- @model_validator(mode="before")
- @classmethod
- def validate_schema(cls, values: Dict) -> Any:
- schema = values["pydantic_schema"]
- if "args_only" not in values:
- values["args_only"] = isinstance(schema, type) and issubclass(
- schema, BaseModel
- )
- elif values["args_only"] and isinstance(schema, Dict):
- raise ValueError(
- "If multiple pydantic schemas are provided then args_only should be"
- " False."
- )
- return values
-
- def parse_result(self, result: List[Generation], *, partial: bool = False) -> Any:
- _result = super().parse_result(result)
- if self.args_only:
- pydantic_args = self.pydantic_schema.parse_raw(_result) # type: ignore[union-attr]
- else:
- fn_name = _result["name"]
- _args = _result["arguments"]
- pydantic_args = self.pydantic_schema[fn_name].parse_raw(_args) # type: ignore[index]
- return pydantic_args
-
-
-class PydanticAttrOutputFunctionsParser(PydanticOutputFunctionsParser):
- """Parse an output as an attribute of a pydantic object."""
-
- attr_name: str
- """The name of the attribute to return."""
-
- def parse_result(self, result: List[Generation], *, partial: bool = False) -> Any:
- result = super().parse_result(result)
- return getattr(result, self.attr_name)
diff --git a/libs/community/langchain_community/output_parsers/rail_parser.py b/libs/community/langchain_community/output_parsers/rail_parser.py
deleted file mode 100644
index f0cabc13eb..0000000000
--- a/libs/community/langchain_community/output_parsers/rail_parser.py
+++ /dev/null
@@ -1,109 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Callable, Dict, Optional
-
-from langchain_core.output_parsers import BaseOutputParser
-
-
-class GuardrailsOutputParser(BaseOutputParser):
- """Parse the output of an LLM call using Guardrails."""
-
- guard: Any
- """The Guardrails object."""
- api: Optional[Callable]
- """The LLM API passed to Guardrails during parsing. An example is `openai.completions.create`.""" # noqa: E501
- args: Any
- """Positional arguments to pass to the above LLM API callable."""
- kwargs: Any
- """Keyword arguments to pass to the above LLM API callable."""
-
- @property
- def _type(self) -> str:
- return "guardrails"
-
- @classmethod
- def from_rail(
- cls,
- rail_file: str,
- num_reasks: int = 1,
- api: Optional[Callable] = None,
- *args: Any,
- **kwargs: Any,
- ) -> GuardrailsOutputParser:
- """Create a GuardrailsOutputParser from a rail file.
-
- Args:
- rail_file: a rail file.
- num_reasks: number of times to re-ask the question.
- api: the API to use for the Guardrails object.
- *args: The arguments to pass to the API
- **kwargs: The keyword arguments to pass to the API.
-
- Returns:
- GuardrailsOutputParser
- """
- try:
- from guardrails import Guard
- except ImportError:
- raise ImportError(
- "guardrails-ai package not installed. "
- "Install it by running `pip install guardrails-ai`."
- )
- return cls(
- guard=Guard.from_rail(rail_file, num_reasks=num_reasks),
- api=api,
- args=args,
- kwargs=kwargs,
- )
-
- @classmethod
- def from_rail_string(
- cls,
- rail_str: str,
- num_reasks: int = 1,
- api: Optional[Callable] = None,
- *args: Any,
- **kwargs: Any,
- ) -> GuardrailsOutputParser:
- try:
- from guardrails import Guard
- except ImportError:
- raise ImportError(
- "guardrails-ai package not installed. "
- "Install it by running `pip install guardrails-ai`."
- )
- return cls(
- guard=Guard.from_rail_string(rail_str, num_reasks=num_reasks),
- api=api,
- args=args,
- kwargs=kwargs,
- )
-
- @classmethod
- def from_pydantic(
- cls,
- output_class: Any,
- num_reasks: int = 1,
- api: Optional[Callable] = None,
- *args: Any,
- **kwargs: Any,
- ) -> GuardrailsOutputParser:
- try:
- from guardrails import Guard
- except ImportError:
- raise ImportError(
- "guardrails-ai package not installed. "
- "Install it by running `pip install guardrails-ai`."
- )
- return cls(
- guard=Guard.from_pydantic(output_class, "", num_reasks=num_reasks),
- api=api,
- args=args,
- kwargs=kwargs,
- )
-
- def get_format_instructions(self) -> str:
- return self.guard.raw_prompt.format_instructions
-
- def parse(self, text: str) -> Dict:
- return self.guard.parse(text, llm_api=self.api, *self.args, **self.kwargs)
diff --git a/libs/community/langchain_community/py.typed b/libs/community/langchain_community/py.typed
deleted file mode 100644
index e69de29bb2..0000000000
diff --git a/libs/community/langchain_community/query_constructors/__init__.py b/libs/community/langchain_community/query_constructors/__init__.py
deleted file mode 100644
index e69de29bb2..0000000000
diff --git a/libs/community/langchain_community/query_constructors/astradb.py b/libs/community/langchain_community/query_constructors/astradb.py
deleted file mode 100644
index ea5be5e18f..0000000000
--- a/libs/community/langchain_community/query_constructors/astradb.py
+++ /dev/null
@@ -1,71 +0,0 @@
-"""Logic for converting internal query language to a valid AstraDB query."""
-
-from typing import Dict, Tuple, Union
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-MULTIPLE_ARITY_COMPARATORS = [Comparator.IN, Comparator.NIN]
-
-
-class AstraDBTranslator(Visitor):
- """Translate AstraDB internal query language elements to valid filters."""
-
- """Subset of allowed logical comparators."""
- allowed_comparators = [
- Comparator.EQ,
- Comparator.NE,
- Comparator.GT,
- Comparator.GTE,
- Comparator.LT,
- Comparator.LTE,
- Comparator.IN,
- Comparator.NIN,
- ]
-
- """Subset of allowed logical operators."""
- allowed_operators = [Operator.AND, Operator.OR]
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- map_dict = {
- Operator.AND: "$and",
- Operator.OR: "$or",
- Comparator.EQ: "$eq",
- Comparator.NE: "$ne",
- Comparator.GTE: "$gte",
- Comparator.LTE: "$lte",
- Comparator.LT: "$lt",
- Comparator.GT: "$gt",
- Comparator.IN: "$in",
- Comparator.NIN: "$nin",
- }
- return map_dict[func]
-
- def visit_operation(self, operation: Operation) -> Dict:
- args = [arg.accept(self) for arg in operation.arguments]
- return {self._format_func(operation.operator): args}
-
- def visit_comparison(self, comparison: Comparison) -> Dict:
- if comparison.comparator in MULTIPLE_ARITY_COMPARATORS and not isinstance(
- comparison.value, list
- ):
- comparison.value = [comparison.value]
-
- comparator = self._format_func(comparison.comparator)
- return {comparison.attribute: {comparator: comparison.value}}
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"filter": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/chroma.py b/libs/community/langchain_community/query_constructors/chroma.py
deleted file mode 100644
index 6f766e7e13..0000000000
--- a/libs/community/langchain_community/query_constructors/chroma.py
+++ /dev/null
@@ -1,50 +0,0 @@
-from typing import Dict, Tuple, Union
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-
-class ChromaTranslator(Visitor):
- """Translate `Chroma` internal query language elements to valid filters."""
-
- allowed_operators = [Operator.AND, Operator.OR]
- """Subset of allowed logical operators."""
- allowed_comparators = [
- Comparator.EQ,
- Comparator.NE,
- Comparator.GT,
- Comparator.GTE,
- Comparator.LT,
- Comparator.LTE,
- ]
- """Subset of allowed logical comparators."""
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- return f"${func.value}"
-
- def visit_operation(self, operation: Operation) -> Dict:
- args = [arg.accept(self) for arg in operation.arguments]
- return {self._format_func(operation.operator): args}
-
- def visit_comparison(self, comparison: Comparison) -> Dict:
- return {
- comparison.attribute: {
- self._format_func(comparison.comparator): comparison.value
- }
- }
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"filter": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/dashvector.py b/libs/community/langchain_community/query_constructors/dashvector.py
deleted file mode 100644
index 65a48d3a81..0000000000
--- a/libs/community/langchain_community/query_constructors/dashvector.py
+++ /dev/null
@@ -1,65 +0,0 @@
-"""Logic for converting internal query language to a valid DashVector query."""
-
-from typing import Tuple, Union
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-
-class DashvectorTranslator(Visitor):
- """Logic for converting internal query language elements to valid filters."""
-
- allowed_operators = [Operator.AND, Operator.OR]
- allowed_comparators = [
- Comparator.EQ,
- Comparator.GT,
- Comparator.GTE,
- Comparator.LT,
- Comparator.LTE,
- Comparator.LIKE,
- ]
-
- map_dict = {
- Operator.AND: " AND ",
- Operator.OR: " OR ",
- Comparator.EQ: " = ",
- Comparator.GT: " > ",
- Comparator.GTE: " >= ",
- Comparator.LT: " < ",
- Comparator.LTE: " <= ",
- Comparator.LIKE: " LIKE ",
- }
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- return self.map_dict[func]
-
- def visit_operation(self, operation: Operation) -> str:
- args = [arg.accept(self) for arg in operation.arguments]
- return self._format_func(operation.operator).join(args)
-
- def visit_comparison(self, comparison: Comparison) -> str:
- value = comparison.value
- if isinstance(value, str):
- if comparison.comparator == Comparator.LIKE:
- value = f"'%{value}%'"
- else:
- value = f"'{value}'"
- return (
- f"{comparison.attribute}{self._format_func(comparison.comparator)}{value}"
- )
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"filter": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/databricks_vector_search.py b/libs/community/langchain_community/query_constructors/databricks_vector_search.py
deleted file mode 100644
index f79a690c9b..0000000000
--- a/libs/community/langchain_community/query_constructors/databricks_vector_search.py
+++ /dev/null
@@ -1,94 +0,0 @@
-from collections import ChainMap
-from itertools import chain
-from typing import Dict, Tuple
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-_COMPARATOR_TO_SYMBOL = {
- Comparator.EQ: "",
- Comparator.GT: " >",
- Comparator.GTE: " >=",
- Comparator.LT: " <",
- Comparator.LTE: " <=",
- Comparator.IN: "",
- Comparator.LIKE: " LIKE",
-}
-
-
-class DatabricksVectorSearchTranslator(Visitor):
- """Translate `Databricks vector search` internal query language elements to
- valid filters."""
-
- """Subset of allowed logical operators."""
- allowed_operators = [Operator.AND, Operator.NOT, Operator.OR]
-
- """Subset of allowed logical comparators."""
- allowed_comparators = [
- Comparator.EQ,
- Comparator.GT,
- Comparator.GTE,
- Comparator.LT,
- Comparator.LTE,
- Comparator.IN,
- Comparator.LIKE,
- ]
-
- def _visit_and_operation(self, operation: Operation) -> Dict:
- return dict(ChainMap(*[arg.accept(self) for arg in operation.arguments]))
-
- def _visit_or_operation(self, operation: Operation) -> Dict:
- filter_args = [arg.accept(self) for arg in operation.arguments]
- flattened_args = list(
- chain.from_iterable(filter_arg.items() for filter_arg in filter_args)
- )
- return {
- " OR ".join(key for key, _ in flattened_args): [
- value for _, value in flattened_args
- ]
- }
-
- def _visit_not_operation(self, operation: Operation) -> Dict:
- if len(operation.arguments) > 1:
- raise ValueError(
- f'"{operation.operator.value}" can have only one argument '
- f"in Databricks vector search"
- )
- filter_arg = operation.arguments[0].accept(self)
- return {
- f"{colum_with_bool_expression} NOT": value
- for colum_with_bool_expression, value in filter_arg.items()
- }
-
- def visit_operation(self, operation: Operation) -> Dict:
- self._validate_func(operation.operator)
- if operation.operator == Operator.AND:
- return self._visit_and_operation(operation)
- elif operation.operator == Operator.OR:
- return self._visit_or_operation(operation)
- elif operation.operator == Operator.NOT:
- return self._visit_not_operation(operation)
- else:
- raise NotImplementedError(
- f'Operator "{operation.operator}" is not supported'
- )
-
- def visit_comparison(self, comparison: Comparison) -> Dict:
- self._validate_func(comparison.comparator)
- comparator_symbol = _COMPARATOR_TO_SYMBOL[comparison.comparator]
- return {f"{comparison.attribute}{comparator_symbol}": comparison.value}
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"filter": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/deeplake.py b/libs/community/langchain_community/query_constructors/deeplake.py
deleted file mode 100644
index d339eb0cdc..0000000000
--- a/libs/community/langchain_community/query_constructors/deeplake.py
+++ /dev/null
@@ -1,89 +0,0 @@
-"""Logic for converting internal query language to a valid Chroma query."""
-
-from typing import Tuple, Union
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-COMPARATOR_TO_TQL = {
- Comparator.EQ: "==",
- Comparator.GT: ">",
- Comparator.GTE: ">=",
- Comparator.LT: "<",
- Comparator.LTE: "<=",
-}
-
-
-OPERATOR_TO_TQL = {
- Operator.AND: "and",
- Operator.OR: "or",
- Operator.NOT: "NOT",
-}
-
-
-def can_cast_to_float(string: str) -> bool:
- """Check if a string can be cast to a float."""
- try:
- float(string)
- return True
- except ValueError:
- return False
-
-
-class DeepLakeTranslator(Visitor):
- """Translate `DeepLake` internal query language elements to valid filters."""
-
- allowed_operators = [Operator.AND, Operator.OR, Operator.NOT]
- """Subset of allowed logical operators."""
- allowed_comparators = [
- Comparator.EQ,
- Comparator.GT,
- Comparator.GTE,
- Comparator.LT,
- Comparator.LTE,
- ]
- """Subset of allowed logical comparators."""
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- if isinstance(func, Operator):
- value = OPERATOR_TO_TQL[func.value] # type: ignore[index]
- elif isinstance(func, Comparator):
- value = COMPARATOR_TO_TQL[func.value] # type: ignore[index]
- return f"{value}"
-
- def visit_operation(self, operation: Operation) -> str:
- args = [arg.accept(self) for arg in operation.arguments]
- operator = self._format_func(operation.operator)
- return "(" + (" " + operator + " ").join(args) + ")"
-
- def visit_comparison(self, comparison: Comparison) -> str:
- comparator = self._format_func(comparison.comparator)
- values = comparison.value
- if isinstance(values, list):
- tql = []
- for value in values:
- comparison.value = value
- tql.append(self.visit_comparison(comparison))
-
- return "(" + (" or ").join(tql) + ")"
-
- if not can_cast_to_float(comparison.value):
- values = f"'{values}'"
- return f"metadata['{comparison.attribute}'] {comparator} {values}"
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- tqL = f"SELECT * WHERE {structured_query.filter.accept(self)}"
- kwargs = {"tql": tqL}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/dingo.py b/libs/community/langchain_community/query_constructors/dingo.py
deleted file mode 100644
index 6c2402f65c..0000000000
--- a/libs/community/langchain_community/query_constructors/dingo.py
+++ /dev/null
@@ -1,49 +0,0 @@
-from typing import Tuple, Union
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-
-class DingoDBTranslator(Visitor):
- """Translate `DingoDB` internal query language elements to valid filters."""
-
- allowed_comparators = (
- Comparator.EQ,
- Comparator.NE,
- Comparator.LT,
- Comparator.LTE,
- Comparator.GT,
- Comparator.GTE,
- )
- """Subset of allowed logical comparators."""
- allowed_operators = (Operator.AND, Operator.OR)
- """Subset of allowed logical operators."""
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- return f"${func.value}"
-
- def visit_operation(self, operation: Operation) -> Operation:
- return operation
-
- def visit_comparison(self, comparison: Comparison) -> Comparison:
- return comparison
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {
- "search_params": {
- "langchain_expr": structured_query.filter.accept(self)
- }
- }
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/elasticsearch.py b/libs/community/langchain_community/query_constructors/elasticsearch.py
deleted file mode 100644
index d07c284b12..0000000000
--- a/libs/community/langchain_community/query_constructors/elasticsearch.py
+++ /dev/null
@@ -1,100 +0,0 @@
-from typing import Dict, Tuple, Union
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-
-class ElasticsearchTranslator(Visitor):
- """Translate `Elasticsearch` internal query language elements to valid filters."""
-
- allowed_comparators = [
- Comparator.EQ,
- Comparator.GT,
- Comparator.GTE,
- Comparator.LT,
- Comparator.LTE,
- Comparator.CONTAIN,
- Comparator.LIKE,
- ]
- """Subset of allowed logical comparators."""
-
- allowed_operators = [Operator.AND, Operator.OR, Operator.NOT]
- """Subset of allowed logical operators."""
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- map_dict = {
- Operator.OR: "should",
- Operator.NOT: "must_not",
- Operator.AND: "must",
- Comparator.EQ: "term",
- Comparator.GT: "gt",
- Comparator.GTE: "gte",
- Comparator.LT: "lt",
- Comparator.LTE: "lte",
- Comparator.CONTAIN: "match",
- Comparator.LIKE: "match",
- }
- return map_dict[func]
-
- def visit_operation(self, operation: Operation) -> Dict:
- args = [arg.accept(self) for arg in operation.arguments]
-
- return {"bool": {self._format_func(operation.operator): args}}
-
- def visit_comparison(self, comparison: Comparison) -> Dict:
- # ElasticsearchStore filters require to target
- # the metadata object field
- field = f"metadata.{comparison.attribute}"
-
- is_range_comparator = comparison.comparator in [
- Comparator.GT,
- Comparator.GTE,
- Comparator.LT,
- Comparator.LTE,
- ]
-
- if is_range_comparator:
- value = comparison.value
- if isinstance(comparison.value, dict) and "date" in comparison.value:
- value = comparison.value["date"]
- return {"range": {field: {self._format_func(comparison.comparator): value}}}
-
- if comparison.comparator == Comparator.CONTAIN:
- return {
- self._format_func(comparison.comparator): {
- field: {"query": comparison.value}
- }
- }
-
- if comparison.comparator == Comparator.LIKE:
- return {
- self._format_func(comparison.comparator): {
- field: {"query": comparison.value, "fuzziness": "AUTO"}
- }
- }
-
- # we assume that if the value is a string,
- # we want to use the keyword field
- field = f"{field}.keyword" if isinstance(comparison.value, str) else field
-
- if isinstance(comparison.value, dict):
- if "date" in comparison.value:
- comparison.value = comparison.value["date"]
-
- return {self._format_func(comparison.comparator): {field: comparison.value}}
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"filter": [structured_query.filter.accept(self)]}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/hanavector.py b/libs/community/langchain_community/query_constructors/hanavector.py
deleted file mode 100644
index 7993782073..0000000000
--- a/libs/community/langchain_community/query_constructors/hanavector.py
+++ /dev/null
@@ -1,75 +0,0 @@
-# HANA Translator/query constructor
-from typing import Dict, Tuple, Union
-
-from langchain_core._api import deprecated
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-
-@deprecated(
- since="0.3.23",
- removal="1.0",
- message=(
- "This class is deprecated and will be removed in a future version. "
- "Please use query_constructors.HanaTranslator from the "
- "langchain_hana package instead. "
- "See https://github.com/SAP/langchain-integration-for-sap-hana-cloud "
- "for details."
- ),
- alternative="from langchain_hana.query_constructors import HanaTranslator;",
- pending=False,
-)
-class HanaTranslator(Visitor):
- """
- **DEPRECATED**: This class is deprecated and will no longer be maintained.
- Please use query_constructors.HanaTranslator from the langchain_hana
- package instead. It offers an improved implementation and full support.
-
- Translate internal query language elements to valid filters params for
- HANA vectorstore.
- """
-
- allowed_operators = [Operator.AND, Operator.OR]
- """Subset of allowed logical operators."""
- allowed_comparators = [
- Comparator.EQ,
- Comparator.NE,
- Comparator.GT,
- Comparator.LT,
- Comparator.GTE,
- Comparator.LTE,
- Comparator.IN,
- Comparator.NIN,
- # Comparator.CONTAIN,
- Comparator.LIKE,
- ]
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- return f"${func.value}"
-
- def visit_operation(self, operation: Operation) -> Dict:
- args = [arg.accept(self) for arg in operation.arguments]
- return {self._format_func(operation.operator): args}
-
- def visit_comparison(self, comparison: Comparison) -> Dict:
- return {
- comparison.attribute: {
- self._format_func(comparison.comparator): comparison.value
- }
- }
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"filter": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/milvus.py b/libs/community/langchain_community/query_constructors/milvus.py
deleted file mode 100644
index a9c2d6f89a..0000000000
--- a/libs/community/langchain_community/query_constructors/milvus.py
+++ /dev/null
@@ -1,104 +0,0 @@
-"""Logic for converting internal query language to a valid Milvus query."""
-
-from typing import Tuple, Union
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-COMPARATOR_TO_BER = {
- Comparator.EQ: "==",
- Comparator.GT: ">",
- Comparator.GTE: ">=",
- Comparator.LT: "<",
- Comparator.LTE: "<=",
- Comparator.IN: "in",
- Comparator.LIKE: "like",
-}
-
-UNARY_OPERATORS = [Operator.NOT]
-
-
-def process_value(value: Union[int, float, str], comparator: Comparator) -> str:
- """Convert a value to a string and add double quotes if it is a string.
-
- It required for comparators involving strings.
-
- Args:
- value: The value to convert.
- comparator: The comparator.
-
- Returns:
- The converted value as a string.
- """
- #
- if isinstance(value, str):
- if comparator is Comparator.LIKE:
- # If the comparator is LIKE, add a percent sign after it for prefix matching
- # and add double quotes
- return f'"{value}%"'
- else:
- # If the value is already a string, add double quotes
- return f'"{value}"'
- else:
- # If the value is not a string, convert it to a string without double quotes
- return str(value)
-
-
-class MilvusTranslator(Visitor):
- """Translate Milvus internal query language elements to valid filters."""
-
- """Subset of allowed logical operators."""
- allowed_operators = [Operator.AND, Operator.NOT, Operator.OR]
-
- """Subset of allowed logical comparators."""
- allowed_comparators = [
- Comparator.EQ,
- Comparator.GT,
- Comparator.GTE,
- Comparator.LT,
- Comparator.LTE,
- Comparator.IN,
- Comparator.LIKE,
- ]
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- value = func.value
- if isinstance(func, Comparator):
- value = COMPARATOR_TO_BER[func]
- return f"{value}"
-
- def visit_operation(self, operation: Operation) -> str:
- if operation.operator in UNARY_OPERATORS and len(operation.arguments) == 1:
- operator = self._format_func(operation.operator)
- return operator + "(" + operation.arguments[0].accept(self) + ")"
- elif operation.operator in UNARY_OPERATORS:
- raise ValueError(
- f'"{operation.operator.value}" can have only one argument in Milvus'
- )
- else:
- args = [arg.accept(self) for arg in operation.arguments]
- operator = self._format_func(operation.operator)
- return "(" + (" " + operator + " ").join(args) + ")"
-
- def visit_comparison(self, comparison: Comparison) -> str:
- comparator = self._format_func(comparison.comparator)
- processed_value = process_value(comparison.value, comparison.comparator)
- attribute = comparison.attribute
-
- return "( " + attribute + " " + comparator + " " + processed_value + " )"
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"expr": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/mongodb_atlas.py b/libs/community/langchain_community/query_constructors/mongodb_atlas.py
deleted file mode 100644
index 1af74fe2b4..0000000000
--- a/libs/community/langchain_community/query_constructors/mongodb_atlas.py
+++ /dev/null
@@ -1,75 +0,0 @@
-"""Logic for converting internal query language to a valid MongoDB Atlas query."""
-
-from typing import Dict, Tuple, Union
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-MULTIPLE_ARITY_COMPARATORS = [Comparator.IN, Comparator.NIN]
-
-
-class MongoDBAtlasTranslator(Visitor):
- """Translate Mongo internal query language elements to valid filters."""
-
- """Subset of allowed logical comparators."""
- allowed_comparators = [
- Comparator.EQ,
- Comparator.NE,
- Comparator.GT,
- Comparator.GTE,
- Comparator.LT,
- Comparator.LTE,
- Comparator.IN,
- Comparator.NIN,
- ]
-
- """Subset of allowed logical operators."""
- allowed_operators = [Operator.AND, Operator.OR]
-
- ## Convert a operator or a comparator to Mongo Query Format
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- map_dict = {
- Operator.AND: "$and",
- Operator.OR: "$or",
- Comparator.EQ: "$eq",
- Comparator.NE: "$ne",
- Comparator.GTE: "$gte",
- Comparator.LTE: "$lte",
- Comparator.LT: "$lt",
- Comparator.GT: "$gt",
- Comparator.IN: "$in",
- Comparator.NIN: "$nin",
- }
- return map_dict[func]
-
- def visit_operation(self, operation: Operation) -> Dict:
- args = [arg.accept(self) for arg in operation.arguments]
- return {self._format_func(operation.operator): args}
-
- def visit_comparison(self, comparison: Comparison) -> Dict:
- if comparison.comparator in MULTIPLE_ARITY_COMPARATORS and not isinstance(
- comparison.value, list
- ):
- comparison.value = [comparison.value]
-
- comparator = self._format_func(comparison.comparator)
-
- attribute = comparison.attribute
-
- return {attribute: {comparator: comparison.value}}
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"pre_filter": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/myscale.py b/libs/community/langchain_community/query_constructors/myscale.py
deleted file mode 100644
index 50a74c568b..0000000000
--- a/libs/community/langchain_community/query_constructors/myscale.py
+++ /dev/null
@@ -1,125 +0,0 @@
-import re
-from typing import Any, Callable, Dict, Tuple
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-
-def _DEFAULT_COMPOSER(op_name: str) -> Callable:
- """
- Default composer for logical operators.
-
- Args:
- op_name: Name of the operator.
-
- Returns:
- Callable that takes a list of arguments and returns a string.
- """
-
- def f(*args: Any) -> str:
- args_: map[str] = map(str, args)
- return f" {op_name} ".join(args_)
-
- return f
-
-
-def _FUNCTION_COMPOSER(op_name: str) -> Callable:
- """
- Composer for functions.
-
- Args:
- op_name: Name of the function.
-
- Returns:
- Callable that takes a list of arguments and returns a string.
- """
-
- def f(*args: Any) -> str:
- args_: map[str] = map(str, args)
- return f"{op_name}({','.join(args_)})"
-
- return f
-
-
-class MyScaleTranslator(Visitor):
- """Translate `MyScale` internal query language elements to valid filters."""
-
- allowed_operators = [Operator.AND, Operator.OR, Operator.NOT]
- """Subset of allowed logical operators."""
-
- allowed_comparators = [
- Comparator.EQ,
- Comparator.GT,
- Comparator.GTE,
- Comparator.LT,
- Comparator.LTE,
- Comparator.CONTAIN,
- Comparator.LIKE,
- ]
-
- map_dict = {
- Operator.AND: _DEFAULT_COMPOSER("AND"),
- Operator.OR: _DEFAULT_COMPOSER("OR"),
- Operator.NOT: _DEFAULT_COMPOSER("NOT"),
- Comparator.EQ: _DEFAULT_COMPOSER("="),
- Comparator.GT: _DEFAULT_COMPOSER(">"),
- Comparator.GTE: _DEFAULT_COMPOSER(">="),
- Comparator.LT: _DEFAULT_COMPOSER("<"),
- Comparator.LTE: _DEFAULT_COMPOSER("<="),
- Comparator.CONTAIN: _FUNCTION_COMPOSER("has"),
- Comparator.LIKE: _DEFAULT_COMPOSER("ILIKE"),
- }
-
- def __init__(self, metadata_key: str = "metadata") -> None:
- super().__init__()
- self.metadata_key = metadata_key
-
- def visit_operation(self, operation: Operation) -> Dict:
- args = [arg.accept(self) for arg in operation.arguments]
- func = operation.operator
- self._validate_func(func)
- return self.map_dict[func](*args)
-
- def visit_comparison(self, comparison: Comparison) -> Dict:
- regex = r"\((.*?)\)"
- matched = re.search(r"\(\w+\)", comparison.attribute)
-
- # If arbitrary function is applied to an attribute
- if matched:
- attr = re.sub(
- regex,
- f"({self.metadata_key}.{matched.group(0)[1:-1]})",
- comparison.attribute,
- )
- else:
- attr = f"{self.metadata_key}.{comparison.attribute}"
- value = comparison.value
- comp = comparison.comparator
-
- value = f"'{value}'" if isinstance(value, str) else value
-
- # convert timestamp for datetime objects
- if isinstance(value, dict) and value.get("type") == "date":
- attr = f"parseDateTime32BestEffort({attr})"
- value = f"parseDateTime32BestEffort('{value['date']}')"
-
- # string pattern match
- if comp is Comparator.LIKE:
- value = f"'%{value[1:-1]}%'"
- return self.map_dict[comp](attr, value)
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- print(structured_query) # noqa: T201
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"where_str": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/neo4j.py b/libs/community/langchain_community/query_constructors/neo4j.py
deleted file mode 100644
index 2ce1de136f..0000000000
--- a/libs/community/langchain_community/query_constructors/neo4j.py
+++ /dev/null
@@ -1,66 +0,0 @@
-from typing import Dict, Tuple, Union
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-
-@deprecated(
- since="0.3.8",
- removal="1.0",
- alternative_import="langchain_neo4j.query_constructors.neo4j.Neo4jTranslator",
-)
-class Neo4jTranslator(Visitor):
- """Translate `Neo4j` internal query language elements to valid filters."""
-
- allowed_operators = [Operator.AND, Operator.OR]
- """Subset of allowed logical operators."""
-
- allowed_comparators = [
- Comparator.EQ,
- Comparator.NE,
- Comparator.GTE,
- Comparator.LTE,
- Comparator.LT,
- Comparator.GT,
- ]
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- map_dict = {
- Operator.AND: "$and",
- Operator.OR: "$or",
- Comparator.EQ: "$eq",
- Comparator.NE: "$ne",
- Comparator.GTE: "$gte",
- Comparator.LTE: "$lte",
- Comparator.LT: "$lt",
- Comparator.GT: "$gt",
- }
- return map_dict[func]
-
- def visit_operation(self, operation: Operation) -> Dict:
- args = [arg.accept(self) for arg in operation.arguments]
- return {self._format_func(operation.operator): args}
-
- def visit_comparison(self, comparison: Comparison) -> Dict:
- return {
- comparison.attribute: {
- self._format_func(comparison.comparator): comparison.value
- }
- }
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"filter": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/opensearch.py b/libs/community/langchain_community/query_constructors/opensearch.py
deleted file mode 100644
index 8b5f23a80c..0000000000
--- a/libs/community/langchain_community/query_constructors/opensearch.py
+++ /dev/null
@@ -1,104 +0,0 @@
-from typing import Dict, Tuple, Union
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-
-class OpenSearchTranslator(Visitor):
- """Translate `OpenSearch` internal query domain-specific
- language elements to valid filters."""
-
- allowed_comparators = [
- Comparator.EQ,
- Comparator.LT,
- Comparator.LTE,
- Comparator.GT,
- Comparator.GTE,
- Comparator.CONTAIN,
- Comparator.LIKE,
- ]
- """Subset of allowed logical comparators."""
-
- allowed_operators = [Operator.AND, Operator.OR, Operator.NOT]
- """Subset of allowed logical operators."""
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- comp_operator_map = {
- Comparator.EQ: "term",
- Comparator.LT: "lt",
- Comparator.LTE: "lte",
- Comparator.GT: "gt",
- Comparator.GTE: "gte",
- Comparator.CONTAIN: "wildcard",
- Comparator.LIKE: "fuzzy",
- Operator.AND: "must",
- Operator.OR: "should",
- Operator.NOT: "must_not",
- }
- return comp_operator_map[func]
-
- def visit_operation(self, operation: Operation) -> Dict:
- args = [arg.accept(self) for arg in operation.arguments]
-
- return {"bool": {self._format_func(operation.operator): args}}
-
- def visit_comparison(self, comparison: Comparison) -> Dict:
- field = f"metadata.{comparison.attribute}"
-
- if comparison.comparator in [
- Comparator.LT,
- Comparator.LTE,
- Comparator.GT,
- Comparator.GTE,
- ]:
- if isinstance(comparison.value, dict):
- if "date" in comparison.value:
- return {
- "range": {
- field: {
- self._format_func(
- comparison.comparator
- ): comparison.value["date"]
- }
- }
- }
- else:
- return {
- "range": {
- field: {
- self._format_func(comparison.comparator): comparison.value
- }
- }
- }
-
- if comparison.comparator == Comparator.LIKE:
- return {
- self._format_func(comparison.comparator): {
- field: {"value": comparison.value}
- }
- }
-
- field = f"{field}.keyword" if isinstance(comparison.value, str) else field
-
- if isinstance(comparison.value, dict):
- if "date" in comparison.value:
- comparison.value = comparison.value["date"]
-
- return {self._format_func(comparison.comparator): {field: comparison.value}}
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"filter": structured_query.filter.accept(self)}
-
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/pgvector.py b/libs/community/langchain_community/query_constructors/pgvector.py
deleted file mode 100644
index 5fea65b01c..0000000000
--- a/libs/community/langchain_community/query_constructors/pgvector.py
+++ /dev/null
@@ -1,52 +0,0 @@
-from typing import Dict, Tuple, Union
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-
-class PGVectorTranslator(Visitor):
- """Translate `PGVector` internal query language elements to valid filters."""
-
- allowed_operators = [Operator.AND, Operator.OR]
- """Subset of allowed logical operators."""
- allowed_comparators = [
- Comparator.EQ,
- Comparator.NE,
- Comparator.GT,
- Comparator.LT,
- Comparator.IN,
- Comparator.NIN,
- Comparator.CONTAIN,
- Comparator.LIKE,
- ]
- """Subset of allowed logical comparators."""
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- return f"{func.value}"
-
- def visit_operation(self, operation: Operation) -> Dict:
- args = [arg.accept(self) for arg in operation.arguments]
- return {self._format_func(operation.operator): args}
-
- def visit_comparison(self, comparison: Comparison) -> Dict:
- return {
- comparison.attribute: {
- self._format_func(comparison.comparator): comparison.value
- }
- }
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"filter": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/pinecone.py b/libs/community/langchain_community/query_constructors/pinecone.py
deleted file mode 100644
index 99c42f393b..0000000000
--- a/libs/community/langchain_community/query_constructors/pinecone.py
+++ /dev/null
@@ -1,57 +0,0 @@
-from typing import Dict, Tuple, Union
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-
-class PineconeTranslator(Visitor):
- """Translate `Pinecone` internal query language elements to valid filters."""
-
- allowed_comparators = (
- Comparator.EQ,
- Comparator.NE,
- Comparator.LT,
- Comparator.LTE,
- Comparator.GT,
- Comparator.GTE,
- Comparator.IN,
- Comparator.NIN,
- )
- """Subset of allowed logical comparators."""
- allowed_operators = (Operator.AND, Operator.OR)
- """Subset of allowed logical operators."""
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- return f"${func.value}"
-
- def visit_operation(self, operation: Operation) -> Dict:
- args = [arg.accept(self) for arg in operation.arguments]
- return {self._format_func(operation.operator): args}
-
- def visit_comparison(self, comparison: Comparison) -> Dict:
- if comparison.comparator in (Comparator.IN, Comparator.NIN) and not isinstance(
- comparison.value, list
- ):
- comparison.value = [comparison.value]
-
- return {
- comparison.attribute: {
- self._format_func(comparison.comparator): comparison.value
- }
- }
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"filter": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/qdrant.py b/libs/community/langchain_community/query_constructors/qdrant.py
deleted file mode 100644
index f4c3298b66..0000000000
--- a/libs/community/langchain_community/query_constructors/qdrant.py
+++ /dev/null
@@ -1,98 +0,0 @@
-from __future__ import annotations
-
-from typing import TYPE_CHECKING, Tuple
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-if TYPE_CHECKING:
- from qdrant_client.http import models as rest
-
-
-class QdrantTranslator(Visitor):
- """Translate `Qdrant` internal query language elements to valid filters."""
-
- allowed_operators = (
- Operator.AND,
- Operator.OR,
- Operator.NOT,
- )
- """Subset of allowed logical operators."""
-
- allowed_comparators = (
- Comparator.EQ,
- Comparator.LT,
- Comparator.LTE,
- Comparator.GT,
- Comparator.GTE,
- Comparator.LIKE,
- )
- """Subset of allowed logical comparators."""
-
- def __init__(self, metadata_key: str):
- self.metadata_key = metadata_key
-
- def visit_operation(self, operation: Operation) -> rest.Filter:
- try:
- from qdrant_client.http import models as rest
- except ImportError as e:
- raise ImportError(
- "Cannot import qdrant_client. Please install with `pip install "
- "qdrant-client`."
- ) from e
-
- args = [arg.accept(self) for arg in operation.arguments]
- operator = {
- Operator.AND: "must",
- Operator.OR: "should",
- Operator.NOT: "must_not",
- }[operation.operator]
- return rest.Filter(**{operator: args})
-
- def visit_comparison(self, comparison: Comparison) -> rest.FieldCondition:
- try:
- from qdrant_client.http import models as rest
- except ImportError as e:
- raise ImportError(
- "Cannot import qdrant_client. Please install with `pip install "
- "qdrant-client`."
- ) from e
-
- self._validate_func(comparison.comparator)
- attribute = self.metadata_key + "." + comparison.attribute
- if comparison.comparator == Comparator.EQ:
- return rest.FieldCondition(
- key=attribute, match=rest.MatchValue(value=comparison.value)
- )
- if comparison.comparator == Comparator.LIKE:
- return rest.FieldCondition(
- key=attribute, match=rest.MatchText(text=comparison.value)
- )
- kwargs = {comparison.comparator.value: comparison.value}
- return rest.FieldCondition(key=attribute, range=rest.Range(**kwargs))
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- try:
- from qdrant_client.http import models as rest
- except ImportError as e:
- raise ImportError(
- "Cannot import qdrant_client. Please install with `pip install "
- "qdrant-client`."
- ) from e
-
- if structured_query.filter is None:
- kwargs = {}
- else:
- filter = structured_query.filter.accept(self)
- if isinstance(filter, rest.FieldCondition):
- filter = rest.Filter(must=[filter])
- kwargs = {"filter": filter}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/redis.py b/libs/community/langchain_community/query_constructors/redis.py
deleted file mode 100644
index e74d1eb199..0000000000
--- a/libs/community/langchain_community/query_constructors/redis.py
+++ /dev/null
@@ -1,103 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Tuple
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-from langchain_community.vectorstores.redis import Redis
-from langchain_community.vectorstores.redis.filters import (
- RedisFilterExpression,
- RedisFilterField,
- RedisFilterOperator,
- RedisNum,
- RedisTag,
- RedisText,
-)
-from langchain_community.vectorstores.redis.schema import RedisModel
-
-_COMPARATOR_TO_BUILTIN_METHOD = {
- Comparator.EQ: "__eq__",
- Comparator.NE: "__ne__",
- Comparator.LT: "__lt__",
- Comparator.GT: "__gt__",
- Comparator.LTE: "__le__",
- Comparator.GTE: "__ge__",
- Comparator.CONTAIN: "__eq__",
- Comparator.LIKE: "__mod__",
-}
-
-
-class RedisTranslator(Visitor):
- """Visitor for translating structured queries to Redis filter expressions."""
-
- allowed_comparators = (
- Comparator.EQ,
- Comparator.NE,
- Comparator.LT,
- Comparator.LTE,
- Comparator.GT,
- Comparator.GTE,
- Comparator.CONTAIN,
- Comparator.LIKE,
- )
- """Subset of allowed logical comparators."""
- allowed_operators = (Operator.AND, Operator.OR)
- """Subset of allowed logical operators."""
-
- def __init__(self, schema: RedisModel) -> None:
- self._schema = schema
-
- def _attribute_to_filter_field(self, attribute: str) -> RedisFilterField:
- if attribute in [tf.name for tf in self._schema.text]:
- return RedisText(attribute)
- elif attribute in [tf.name for tf in self._schema.tag or []]:
- return RedisTag(attribute)
- elif attribute in [tf.name for tf in self._schema.numeric or []]:
- return RedisNum(attribute)
- else:
- raise ValueError(
- f"Invalid attribute {attribute} not in vector store schema. Schema is:"
- f"\n{self._schema.as_dict()}"
- )
-
- def visit_comparison(self, comparison: Comparison) -> RedisFilterExpression:
- filter_field = self._attribute_to_filter_field(comparison.attribute)
- comparison_method = _COMPARATOR_TO_BUILTIN_METHOD[comparison.comparator]
- return getattr(filter_field, comparison_method)(comparison.value)
-
- def visit_operation(self, operation: Operation) -> Any:
- left = operation.arguments[0].accept(self)
- if len(operation.arguments) > 2:
- right = self.visit_operation(
- Operation(
- operator=operation.operator, arguments=operation.arguments[1:]
- )
- )
- else:
- right = operation.arguments[1].accept(self)
- redis_operator = (
- RedisFilterOperator.OR
- if operation.operator == Operator.OR
- else RedisFilterOperator.AND
- )
- return RedisFilterExpression(operator=redis_operator, left=left, right=right)
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"filter": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
-
- @classmethod
- def from_vectorstore(cls, vectorstore: Redis) -> RedisTranslator:
- return cls(vectorstore._schema)
diff --git a/libs/community/langchain_community/query_constructors/supabase.py b/libs/community/langchain_community/query_constructors/supabase.py
deleted file mode 100644
index 8910b3a9d8..0000000000
--- a/libs/community/langchain_community/query_constructors/supabase.py
+++ /dev/null
@@ -1,97 +0,0 @@
-from typing import Any, Dict, Tuple
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-
-class SupabaseVectorTranslator(Visitor):
- """Translate Langchain filters to Supabase PostgREST filters."""
-
- allowed_operators = [Operator.AND, Operator.OR]
- """Subset of allowed logical operators."""
-
- allowed_comparators = [
- Comparator.EQ,
- Comparator.NE,
- Comparator.GT,
- Comparator.GTE,
- Comparator.LT,
- Comparator.LTE,
- Comparator.LIKE,
- ]
- """Subset of allowed logical comparators."""
-
- metadata_column: str = "metadata"
-
- def _map_comparator(self, comparator: Comparator) -> str:
- """
- Maps Langchain comparator to PostgREST comparator:
-
- https://postgrest.org/en/stable/references/api/tables_views.html#operators
- """
- postgrest_comparator = {
- Comparator.EQ: "eq",
- Comparator.NE: "neq",
- Comparator.GT: "gt",
- Comparator.GTE: "gte",
- Comparator.LT: "lt",
- Comparator.LTE: "lte",
- Comparator.LIKE: "like",
- }.get(comparator)
-
- if postgrest_comparator is None:
- raise Exception(
- f"Comparator '{comparator}' is not currently "
- "supported in Supabase Vector"
- )
-
- return postgrest_comparator
-
- def _get_json_operator(self, value: Any) -> str:
- if isinstance(value, str):
- return "->>"
- else:
- return "->"
-
- def visit_operation(self, operation: Operation) -> str:
- args = [arg.accept(self) for arg in operation.arguments]
- return f"{operation.operator.value}({','.join(args)})"
-
- def visit_comparison(self, comparison: Comparison) -> str:
- if isinstance(comparison.value, list):
- return self.visit_operation(
- Operation(
- operator=Operator.AND,
- arguments=[
- Comparison(
- comparator=comparison.comparator,
- attribute=comparison.attribute,
- value=value,
- )
- for value in comparison.value
- ],
- )
- )
-
- return ".".join(
- [
- f"{self.metadata_column}{self._get_json_operator(comparison.value)}{comparison.attribute}",
- f"{self._map_comparator(comparison.comparator)}",
- f"{comparison.value}",
- ]
- )
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, Dict[str, str]]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"postgrest_filter": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/tencentvectordb.py b/libs/community/langchain_community/query_constructors/tencentvectordb.py
deleted file mode 100644
index b1ec31a1a2..0000000000
--- a/libs/community/langchain_community/query_constructors/tencentvectordb.py
+++ /dev/null
@@ -1,116 +0,0 @@
-from __future__ import annotations
-
-from typing import Optional, Sequence, Tuple
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-
-class TencentVectorDBTranslator(Visitor):
- """Translate StructuredQuery to Tencent VectorDB query."""
-
- COMPARATOR_MAP = {
- Comparator.EQ: "=",
- Comparator.NE: "!=",
- Comparator.GT: ">",
- Comparator.GTE: ">=",
- Comparator.LT: "<",
- Comparator.LTE: "<=",
- Comparator.IN: "in",
- Comparator.NIN: "not in",
- }
-
- allowed_comparators: Optional[Sequence[Comparator]] = list(COMPARATOR_MAP.keys())
- allowed_operators: Optional[Sequence[Operator]] = [
- Operator.AND,
- Operator.OR,
- Operator.NOT,
- ]
-
- def __init__(self, meta_keys: Optional[Sequence[str]] = None):
- """Initialize the translator.
-
- Args:
- meta_keys: List of meta keys to be used in the query. Default: [].
- """
- self.meta_keys = meta_keys or []
-
- def visit_operation(self, operation: Operation) -> str:
- """Visit an operation node and return the translated query.
-
- Args:
- operation: Operation node to be visited.
-
- Returns:
- Translated query.
- """
- if operation.operator in (Operator.AND, Operator.OR):
- ret = f" {operation.operator.value} ".join(
- [arg.accept(self) for arg in operation.arguments]
- )
- if operation.operator == Operator.OR:
- ret = f"({ret})"
- return ret
- else:
- return f"not ({operation.arguments[0].accept(self)})"
-
- def visit_comparison(self, comparison: Comparison) -> str:
- """Visit a comparison node and return the translated query.
-
- Args:
- comparison: Comparison node to be visited.
-
- Returns:
- Translated query.
- """
- if self.meta_keys and comparison.attribute not in self.meta_keys:
- raise ValueError(
- f"Expr Filtering found Unsupported attribute: {comparison.attribute}"
- )
-
- if comparison.comparator in self.COMPARATOR_MAP:
- if comparison.comparator in [Comparator.IN, Comparator.NIN]:
- value = map(
- lambda x: f'"{x}"' if isinstance(x, str) else x, comparison.value
- )
- return (
- f"{comparison.attribute}"
- f" {self.COMPARATOR_MAP[comparison.comparator]} "
- f"({', '.join(value)})"
- )
- if isinstance(comparison.value, str):
- return (
- f"{comparison.attribute} "
- f"{self.COMPARATOR_MAP[comparison.comparator]}"
- f' "{comparison.value}"'
- )
- return (
- f"{comparison.attribute}"
- f" {self.COMPARATOR_MAP[comparison.comparator]} "
- f"{comparison.value}"
- )
- else:
- raise ValueError(f"Unsupported comparator {comparison.comparator}")
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- """Visit a structured query node and return the translated query.
-
- Args:
- structured_query: StructuredQuery node to be visited.
-
- Returns:
- Translated query and query kwargs.
- """
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"expr": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/timescalevector.py b/libs/community/langchain_community/query_constructors/timescalevector.py
deleted file mode 100644
index c51718a11e..0000000000
--- a/libs/community/langchain_community/query_constructors/timescalevector.py
+++ /dev/null
@@ -1,84 +0,0 @@
-from __future__ import annotations
-
-from typing import TYPE_CHECKING, Tuple, Union
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-if TYPE_CHECKING:
- from timescale_vector import client
-
-
-class TimescaleVectorTranslator(Visitor):
- """Translate the internal query language elements to valid filters."""
-
- allowed_operators = [Operator.AND, Operator.OR, Operator.NOT]
- """Subset of allowed logical operators."""
-
- allowed_comparators = [
- Comparator.EQ,
- Comparator.GT,
- Comparator.GTE,
- Comparator.LT,
- Comparator.LTE,
- ]
-
- COMPARATOR_MAP = {
- Comparator.EQ: "==",
- Comparator.GT: ">",
- Comparator.GTE: ">=",
- Comparator.LT: "<",
- Comparator.LTE: "<=",
- }
-
- OPERATOR_MAP = {Operator.AND: "AND", Operator.OR: "OR", Operator.NOT: "NOT"}
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- if isinstance(func, Operator):
- value = self.OPERATOR_MAP[func.value] # type: ignore[index]
- elif isinstance(func, Comparator):
- value = self.COMPARATOR_MAP[func.value] # type: ignore[index]
- return f"{value}"
-
- def visit_operation(self, operation: Operation) -> client.Predicates:
- try:
- from timescale_vector import client
- except ImportError as e:
- raise ImportError(
- "Cannot import timescale-vector. Please install with `pip install "
- "timescale-vector`."
- ) from e
- args = [arg.accept(self) for arg in operation.arguments]
- return client.Predicates(*args, operator=self._format_func(operation.operator))
-
- def visit_comparison(self, comparison: Comparison) -> client.Predicates:
- try:
- from timescale_vector import client
- except ImportError as e:
- raise ImportError(
- "Cannot import timescale-vector. Please install with `pip install "
- "timescale-vector`."
- ) from e
- return client.Predicates(
- (
- comparison.attribute,
- self._format_func(comparison.comparator),
- comparison.value,
- )
- )
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"predicates": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/vectara.py b/libs/community/langchain_community/query_constructors/vectara.py
deleted file mode 100644
index 24886a1af9..0000000000
--- a/libs/community/langchain_community/query_constructors/vectara.py
+++ /dev/null
@@ -1,70 +0,0 @@
-from typing import Tuple, Union
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-
-def process_value(value: Union[int, float, str]) -> str:
- """Convert a value to a string and add single quotes if it is a string."""
- if isinstance(value, str):
- return f"'{value}'"
- else:
- return str(value)
-
-
-class VectaraTranslator(Visitor):
- """Translate `Vectara` internal query language elements to valid filters."""
-
- allowed_operators = [Operator.AND, Operator.OR]
- """Subset of allowed logical operators."""
- allowed_comparators = [
- Comparator.EQ,
- Comparator.NE,
- Comparator.GT,
- Comparator.GTE,
- Comparator.LT,
- Comparator.LTE,
- ]
- """Subset of allowed logical comparators."""
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- map_dict = {
- Operator.AND: " and ",
- Operator.OR: " or ",
- Comparator.EQ: "=",
- Comparator.NE: "!=",
- Comparator.GT: ">",
- Comparator.GTE: ">=",
- Comparator.LT: "<",
- Comparator.LTE: "<=",
- }
- self._validate_func(func)
- return map_dict[func]
-
- def visit_operation(self, operation: Operation) -> str:
- args = [arg.accept(self) for arg in operation.arguments]
- operator = self._format_func(operation.operator)
- return "( " + operator.join(args) + " )"
-
- def visit_comparison(self, comparison: Comparison) -> str:
- comparator = self._format_func(comparison.comparator)
- processed_value = process_value(comparison.value)
- attribute = comparison.attribute
- return (
- "( " + "doc." + attribute + " " + comparator + " " + processed_value + " )"
- )
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"filter": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/query_constructors/weaviate.py b/libs/community/langchain_community/query_constructors/weaviate.py
deleted file mode 100644
index 2e5e3e691e..0000000000
--- a/libs/community/langchain_community/query_constructors/weaviate.py
+++ /dev/null
@@ -1,79 +0,0 @@
-from datetime import datetime
-from typing import Dict, Tuple, Union
-
-from langchain_core.structured_query import (
- Comparator,
- Comparison,
- Operation,
- Operator,
- StructuredQuery,
- Visitor,
-)
-
-
-class WeaviateTranslator(Visitor):
- """Translate `Weaviate` internal query language elements to valid filters."""
-
- allowed_operators = [Operator.AND, Operator.OR]
- """Subset of allowed logical operators."""
-
- allowed_comparators = [
- Comparator.EQ,
- Comparator.NE,
- Comparator.GTE,
- Comparator.LTE,
- Comparator.LT,
- Comparator.GT,
- ]
-
- def _format_func(self, func: Union[Operator, Comparator]) -> str:
- self._validate_func(func)
- # https://weaviate.io/developers/weaviate/api/graphql/filters
- map_dict = {
- Operator.AND: "And",
- Operator.OR: "Or",
- Comparator.EQ: "Equal",
- Comparator.NE: "NotEqual",
- Comparator.GTE: "GreaterThanEqual",
- Comparator.LTE: "LessThanEqual",
- Comparator.LT: "LessThan",
- Comparator.GT: "GreaterThan",
- }
- return map_dict[func]
-
- def visit_operation(self, operation: Operation) -> Dict:
- args = [arg.accept(self) for arg in operation.arguments]
- return {"operator": self._format_func(operation.operator), "operands": args}
-
- def visit_comparison(self, comparison: Comparison) -> Dict:
- value_type = "valueText"
- value = comparison.value
- if isinstance(comparison.value, bool):
- value_type = "valueBoolean"
- elif isinstance(comparison.value, float):
- value_type = "valueNumber"
- elif isinstance(comparison.value, int):
- value_type = "valueInt"
- elif (
- isinstance(comparison.value, dict)
- and comparison.value.get("type") == "date"
- ):
- value_type = "valueDate"
- # ISO 8601 timestamp, formatted as RFC3339
- date = datetime.strptime(comparison.value["date"], "%Y-%m-%d")
- value = date.strftime("%Y-%m-%dT%H:%M:%SZ")
- filter = {
- "path": [comparison.attribute],
- "operator": self._format_func(comparison.comparator),
- value_type: value,
- }
- return filter
-
- def visit_structured_query(
- self, structured_query: StructuredQuery
- ) -> Tuple[str, dict]:
- if structured_query.filter is None:
- kwargs = {}
- else:
- kwargs = {"where_filter": structured_query.filter.accept(self)}
- return structured_query.query, kwargs
diff --git a/libs/community/langchain_community/retrievers/__init__.py b/libs/community/langchain_community/retrievers/__init__.py
deleted file mode 100644
index ce4ac731bd..0000000000
--- a/libs/community/langchain_community/retrievers/__init__.py
+++ /dev/null
@@ -1,253 +0,0 @@
-"""**Retriever** class returns Documents given a text **query**.
-
-It is more general than a vector store. A retriever does not need to be able to
-store documents, only to return (or retrieve) it. Vector stores can be used as
-the backbone of a retriever, but there are other types of retrievers as well.
-
-**Class hierarchy:**
-
-.. code-block::
-
- BaseRetriever --> Retriever # Examples: ArxivRetriever, MergerRetriever
-
-**Main helpers:**
-
-.. code-block::
-
- Document, Serializable, Callbacks,
- CallbackManagerForRetrieverRun, AsyncCallbackManagerForRetrieverRun
-"""
-
-import importlib
-from typing import TYPE_CHECKING, Any
-
-if TYPE_CHECKING:
- from langchain_community.retrievers.arcee import (
- ArceeRetriever,
- )
- from langchain_community.retrievers.arxiv import (
- ArxivRetriever,
- )
- from langchain_community.retrievers.asknews import (
- AskNewsRetriever,
- )
- from langchain_community.retrievers.azure_ai_search import (
- AzureAISearchRetriever,
- AzureCognitiveSearchRetriever,
- )
- from langchain_community.retrievers.bedrock import (
- AmazonKnowledgeBasesRetriever,
- )
- from langchain_community.retrievers.bm25 import (
- BM25Retriever,
- )
- from langchain_community.retrievers.breebs import (
- BreebsRetriever,
- )
- from langchain_community.retrievers.chaindesk import (
- ChaindeskRetriever,
- )
- from langchain_community.retrievers.chatgpt_plugin_retriever import (
- ChatGPTPluginRetriever,
- )
- from langchain_community.retrievers.cohere_rag_retriever import (
- CohereRagRetriever,
- )
- from langchain_community.retrievers.docarray import (
- DocArrayRetriever,
- )
- from langchain_community.retrievers.dria_index import (
- DriaRetriever,
- )
- from langchain_community.retrievers.elastic_search_bm25 import (
- ElasticSearchBM25Retriever,
- )
- from langchain_community.retrievers.embedchain import (
- EmbedchainRetriever,
- )
- from langchain_community.retrievers.google_cloud_documentai_warehouse import (
- GoogleDocumentAIWarehouseRetriever,
- )
- from langchain_community.retrievers.google_vertex_ai_search import (
- GoogleCloudEnterpriseSearchRetriever,
- GoogleVertexAIMultiTurnSearchRetriever,
- GoogleVertexAISearchRetriever,
- )
- from langchain_community.retrievers.kay import (
- KayAiRetriever,
- )
- from langchain_community.retrievers.kendra import (
- AmazonKendraRetriever,
- )
- from langchain_community.retrievers.knn import (
- KNNRetriever,
- )
- from langchain_community.retrievers.llama_index import (
- LlamaIndexGraphRetriever,
- LlamaIndexRetriever,
- )
- from langchain_community.retrievers.metal import (
- MetalRetriever,
- )
- from langchain_community.retrievers.milvus import (
- MilvusRetriever,
- )
- from langchain_community.retrievers.nanopq import NanoPQRetriever
- from langchain_community.retrievers.needle import NeedleRetriever
- from langchain_community.retrievers.outline import (
- OutlineRetriever,
- )
- from langchain_community.retrievers.pinecone_hybrid_search import (
- PineconeHybridSearchRetriever,
- )
- from langchain_community.retrievers.pubmed import (
- PubMedRetriever,
- )
- from langchain_community.retrievers.qdrant_sparse_vector_retriever import (
- QdrantSparseVectorRetriever,
- )
- from langchain_community.retrievers.rememberizer import (
- RememberizerRetriever,
- )
- from langchain_community.retrievers.remote_retriever import (
- RemoteLangChainRetriever,
- )
- from langchain_community.retrievers.svm import (
- SVMRetriever,
- )
- from langchain_community.retrievers.tavily_search_api import (
- TavilySearchAPIRetriever,
- )
- from langchain_community.retrievers.tfidf import (
- TFIDFRetriever,
- )
- from langchain_community.retrievers.thirdai_neuraldb import NeuralDBRetriever
- from langchain_community.retrievers.vespa_retriever import (
- VespaRetriever,
- )
- from langchain_community.retrievers.weaviate_hybrid_search import (
- WeaviateHybridSearchRetriever,
- )
- from langchain_community.retrievers.web_research import WebResearchRetriever
- from langchain_community.retrievers.wikipedia import (
- WikipediaRetriever,
- )
- from langchain_community.retrievers.you import (
- YouRetriever,
- )
- from langchain_community.retrievers.zep import (
- ZepRetriever,
- )
- from langchain_community.retrievers.zep_cloud import (
- ZepCloudRetriever,
- )
- from langchain_community.retrievers.zilliz import (
- ZillizRetriever,
- )
-
-
-_module_lookup = {
- "AmazonKendraRetriever": "langchain_community.retrievers.kendra",
- "AmazonKnowledgeBasesRetriever": "langchain_community.retrievers.bedrock",
- "ArceeRetriever": "langchain_community.retrievers.arcee",
- "ArxivRetriever": "langchain_community.retrievers.arxiv",
- "AskNewsRetriever": "langchain_community.retrievers.asknews",
- "AzureAISearchRetriever": "langchain_community.retrievers.azure_ai_search",
- "AzureCognitiveSearchRetriever": "langchain_community.retrievers.azure_ai_search",
- "BM25Retriever": "langchain_community.retrievers.bm25",
- "BreebsRetriever": "langchain_community.retrievers.breebs",
- "ChaindeskRetriever": "langchain_community.retrievers.chaindesk",
- "ChatGPTPluginRetriever": "langchain_community.retrievers.chatgpt_plugin_retriever",
- "CohereRagRetriever": "langchain_community.retrievers.cohere_rag_retriever",
- "DocArrayRetriever": "langchain_community.retrievers.docarray",
- "DriaRetriever": "langchain_community.retrievers.dria_index",
- "ElasticSearchBM25Retriever": "langchain_community.retrievers.elastic_search_bm25",
- "EmbedchainRetriever": "langchain_community.retrievers.embedchain",
- "GoogleCloudEnterpriseSearchRetriever": "langchain_community.retrievers.google_vertex_ai_search", # noqa: E501
- "GoogleDocumentAIWarehouseRetriever": "langchain_community.retrievers.google_cloud_documentai_warehouse", # noqa: E501
- "GoogleVertexAIMultiTurnSearchRetriever": "langchain_community.retrievers.google_vertex_ai_search", # noqa: E501
- "GoogleVertexAISearchRetriever": "langchain_community.retrievers.google_vertex_ai_search", # noqa: E501
- "KNNRetriever": "langchain_community.retrievers.knn",
- "KayAiRetriever": "langchain_community.retrievers.kay",
- "LlamaIndexGraphRetriever": "langchain_community.retrievers.llama_index",
- "LlamaIndexRetriever": "langchain_community.retrievers.llama_index",
- "MetalRetriever": "langchain_community.retrievers.metal",
- "MilvusRetriever": "langchain_community.retrievers.milvus",
- "NanoPQRetriever": "langchain_community.retrievers.nanopq",
- "NeedleRetriever": "langchain_community.retrievers.needle",
- "OutlineRetriever": "langchain_community.retrievers.outline",
- "PineconeHybridSearchRetriever": "langchain_community.retrievers.pinecone_hybrid_search", # noqa: E501
- "PubMedRetriever": "langchain_community.retrievers.pubmed",
- "QdrantSparseVectorRetriever": "langchain_community.retrievers.qdrant_sparse_vector_retriever", # noqa: E501
- "RememberizerRetriever": "langchain_community.retrievers.rememberizer",
- "RemoteLangChainRetriever": "langchain_community.retrievers.remote_retriever",
- "SVMRetriever": "langchain_community.retrievers.svm",
- "TFIDFRetriever": "langchain_community.retrievers.tfidf",
- "TavilySearchAPIRetriever": "langchain_community.retrievers.tavily_search_api",
- "VespaRetriever": "langchain_community.retrievers.vespa_retriever",
- "WeaviateHybridSearchRetriever": "langchain_community.retrievers.weaviate_hybrid_search", # noqa: E501
- "WebResearchRetriever": "langchain_community.retrievers.web_research",
- "WikipediaRetriever": "langchain_community.retrievers.wikipedia",
- "YouRetriever": "langchain_community.retrievers.you",
- "ZepRetriever": "langchain_community.retrievers.zep",
- "ZepCloudRetriever": "langchain_community.retrievers.zep_cloud",
- "ZillizRetriever": "langchain_community.retrievers.zilliz",
- "NeuralDBRetriever": "langchain_community.retrievers.thirdai_neuraldb",
-}
-
-
-def __getattr__(name: str) -> Any:
- if name in _module_lookup:
- module = importlib.import_module(_module_lookup[name])
- return getattr(module, name)
- raise AttributeError(f"module {__name__} has no attribute {name}")
-
-
-__all__ = [
- "AmazonKendraRetriever",
- "AmazonKnowledgeBasesRetriever",
- "ArceeRetriever",
- "ArxivRetriever",
- "AskNewsRetriever",
- "AzureAISearchRetriever",
- "AzureCognitiveSearchRetriever",
- "BM25Retriever",
- "BreebsRetriever",
- "ChaindeskRetriever",
- "ChatGPTPluginRetriever",
- "CohereRagRetriever",
- "DocArrayRetriever",
- "DriaRetriever",
- "ElasticSearchBM25Retriever",
- "EmbedchainRetriever",
- "GoogleCloudEnterpriseSearchRetriever",
- "GoogleDocumentAIWarehouseRetriever",
- "GoogleVertexAIMultiTurnSearchRetriever",
- "GoogleVertexAISearchRetriever",
- "KayAiRetriever",
- "KNNRetriever",
- "LlamaIndexGraphRetriever",
- "LlamaIndexRetriever",
- "MetalRetriever",
- "MilvusRetriever",
- "NanoPQRetriever",
- "NeedleRetriever",
- "NeuralDBRetriever",
- "OutlineRetriever",
- "PineconeHybridSearchRetriever",
- "PubMedRetriever",
- "QdrantSparseVectorRetriever",
- "RememberizerRetriever",
- "RemoteLangChainRetriever",
- "SVMRetriever",
- "TavilySearchAPIRetriever",
- "TFIDFRetriever",
- "VespaRetriever",
- "WeaviateHybridSearchRetriever",
- "WebResearchRetriever",
- "WikipediaRetriever",
- "YouRetriever",
- "ZepRetriever",
- "ZepCloudRetriever",
- "ZillizRetriever",
-]
diff --git a/libs/community/langchain_community/retrievers/arcee.py b/libs/community/langchain_community/retrievers/arcee.py
deleted file mode 100644
index ebdc8301cf..0000000000
--- a/libs/community/langchain_community/retrievers/arcee.py
+++ /dev/null
@@ -1,137 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import ConfigDict, SecretStr
-
-from langchain_community.utilities.arcee import ArceeWrapper, DALMFilter
-
-
-class ArceeRetriever(BaseRetriever):
- """Arcee Domain Adapted Language Models (DALMs) retriever.
-
- To use, set the ``ARCEE_API_KEY`` environment variable with your Arcee API key,
- or pass ``arcee_api_key`` as a named parameter.
-
- Example:
- .. code-block:: python
-
- from langchain_community.retrievers import ArceeRetriever
-
- retriever = ArceeRetriever(
- model="DALM-PubMed",
- arcee_api_key="ARCEE-API-KEY"
- )
-
- documents = retriever.invoke("AI-driven music therapy")
- """
-
- _client: Optional[ArceeWrapper] = None #: :meta private:
- """Arcee client."""
-
- arcee_api_key: SecretStr
- """Arcee API Key"""
-
- model: str
- """Arcee DALM name"""
-
- arcee_api_url: str = "https://api.arcee.ai"
- """Arcee API URL"""
-
- arcee_api_version: str = "v2"
- """Arcee API Version"""
-
- arcee_app_url: str = "https://app.arcee.ai"
- """Arcee App URL"""
-
- model_kwargs: Optional[Dict[str, Any]] = None
- """Keyword arguments to pass to the model."""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- def __init__(self, **data: Any) -> None:
- """Initializes private fields."""
-
- super().__init__(**data)
-
- self._client = ArceeWrapper(
- arcee_api_key=self.arcee_api_key.get_secret_value(),
- arcee_api_url=self.arcee_api_url,
- arcee_api_version=self.arcee_api_version,
- model_kwargs=self.model_kwargs,
- model_name=self.model,
- )
-
- self._client.validate_model_training_status()
-
- @pre_init
- def validate_environments(cls, values: Dict) -> Dict:
- """Validate Arcee environment variables."""
-
- # validate env vars
- values["arcee_api_key"] = convert_to_secret_str(
- get_from_dict_or_env(
- values,
- "arcee_api_key",
- "ARCEE_API_KEY",
- )
- )
-
- values["arcee_api_url"] = get_from_dict_or_env(
- values,
- "arcee_api_url",
- "ARCEE_API_URL",
- )
-
- values["arcee_app_url"] = get_from_dict_or_env(
- values,
- "arcee_app_url",
- "ARCEE_APP_URL",
- )
-
- values["arcee_api_version"] = get_from_dict_or_env(
- values,
- "arcee_api_version",
- "ARCEE_API_VERSION",
- )
-
- # validate model kwargs
- if values["model_kwargs"]:
- kw = values["model_kwargs"]
-
- # validate size
- if kw.get("size") is not None:
- if not kw.get("size") >= 0:
- raise ValueError("`size` must not be negative.")
-
- # validate filters
- if kw.get("filters") is not None:
- if not isinstance(kw.get("filters"), List):
- raise ValueError("`filters` must be a list.")
- for f in kw.get("filters"):
- DALMFilter(**f)
-
- return values
-
- def _get_relevant_documents(
- self, query: str, run_manager: CallbackManagerForRetrieverRun, **kwargs: Any
- ) -> List[Document]:
- """Retrieve {size} contexts with your retriever for a given query
-
- Args:
- query: Query to submit to the model
- size: The max number of context results to retrieve.
- Defaults to 3. (Can be less if filters are provided).
- filters: Filters to apply to the context dataset.
- """
-
- try:
- if not self._client:
- raise ValueError("Client is not initialized.")
- return self._client.retrieve(query=query, **kwargs)
- except Exception as e:
- raise ValueError(f"Error while retrieving documents: {e}") from e
diff --git a/libs/community/langchain_community/retrievers/arxiv.py b/libs/community/langchain_community/retrievers/arxiv.py
deleted file mode 100644
index 3d59e949d5..0000000000
--- a/libs/community/langchain_community/retrievers/arxiv.py
+++ /dev/null
@@ -1,92 +0,0 @@
-from typing import List
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-from langchain_community.utilities.arxiv import ArxivAPIWrapper
-
-
-class ArxivRetriever(BaseRetriever, ArxivAPIWrapper):
- """`Arxiv` retriever.
-
- Setup:
- Install ``arxiv``:
-
- .. code-block:: bash
-
- pip install -U arxiv
-
- Key init args:
- load_max_docs: int
- maximum number of documents to load
- get_ful_documents: bool
- whether to return full document text or snippets
-
- Instantiate:
- .. code-block:: python
-
- from langchain_community.retrievers import ArxivRetriever
-
- retriever = ArxivRetriever(
- load_max_docs=2,
- get_ful_documents=True,
- )
-
- Usage:
- .. code-block:: python
-
- docs = retriever.invoke("What is the ImageBind model?")
- docs[0].metadata
-
- .. code-block:: none
-
- {'Entry ID': 'http://arxiv.org/abs/2305.05665v2',
- 'Published': datetime.date(2023, 5, 31),
- 'Title': 'ImageBind: One Embedding Space To Bind Them All',
- 'Authors': 'Rohit Girdhar, Alaaeldin El-Nouby, Zhuang Liu, Mannat Singh, Kalyan Vasudev Alwala, Armand Joulin, Ishan Misra'}
-
- Use within a chain:
- .. code-block:: python
-
- from langchain_core.output_parsers import StrOutputParser
- from langchain_core.prompts import ChatPromptTemplate
- from langchain_core.runnables import RunnablePassthrough
- from langchain_openai import ChatOpenAI
-
- prompt = ChatPromptTemplate.from_template(
- \"\"\"Answer the question based only on the context provided.
-
- Context: {context}
-
- Question: {question}\"\"\"
- )
-
- llm = ChatOpenAI(model="gpt-3.5-turbo-0125")
-
- def format_docs(docs):
- return "\\n\\n".join(doc.page_content for doc in docs)
-
- chain = (
- {"context": retriever | format_docs, "question": RunnablePassthrough()}
- | prompt
- | llm
- | StrOutputParser()
- )
-
- chain.invoke("What is the ImageBind model?")
-
- .. code-block:: none
-
- 'The ImageBind model is an approach to learn a joint embedding across six different modalities - images, text, audio, depth, thermal, and IMU data...'
- """ # noqa: E501
-
- get_full_documents: bool = False
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- if self.get_full_documents:
- return self.load(query=query)
- else:
- return self.get_summaries_as_docs(query)
diff --git a/libs/community/langchain_community/retrievers/asknews.py b/libs/community/langchain_community/retrievers/asknews.py
deleted file mode 100644
index 18a44161d2..0000000000
--- a/libs/community/langchain_community/retrievers/asknews.py
+++ /dev/null
@@ -1,146 +0,0 @@
-import os
-import re
-from typing import Any, Dict, List, Literal, Optional
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForRetrieverRun,
- CallbackManagerForRetrieverRun,
-)
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-
-class AskNewsRetriever(BaseRetriever):
- """AskNews retriever."""
-
- k: int = 10
- offset: int = 0
- start_timestamp: Optional[int] = None
- end_timestamp: Optional[int] = None
- method: Literal["nl", "kw"] = "nl"
- categories: List[
- Literal[
- "All",
- "Business",
- "Crime",
- "Politics",
- "Science",
- "Sports",
- "Technology",
- "Military",
- "Health",
- "Entertainment",
- "Finance",
- "Culture",
- "Climate",
- "Environment",
- "World",
- ]
- ] = ["All"]
- historical: bool = False
- similarity_score_threshold: float = 0.5
- kwargs: Optional[Dict[str, Any]] = {}
- client_id: Optional[str] = None
- client_secret: Optional[str] = None
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- """Get documents relevant to a query.
- Args:
- query: String to find relevant documents for
- run_manager: The callbacks handler to use
- Returns:
- List of relevant documents
- """
- try:
- from asknews_sdk import AskNewsSDK
- except ImportError:
- raise ImportError(
- "AskNews python package not found. "
- "Please install it with `pip install asknews`."
- )
- an_client = AskNewsSDK(
- client_id=self.client_id or os.environ["ASKNEWS_CLIENT_ID"],
- client_secret=self.client_secret or os.environ["ASKNEWS_CLIENT_SECRET"],
- scopes=["news"],
- )
- response = an_client.news.search_news(
- query=query,
- n_articles=self.k,
- start_timestamp=self.start_timestamp,
- end_timestamp=self.end_timestamp,
- method=self.method,
- categories=self.categories,
- historical=self.historical,
- similarity_score_threshold=self.similarity_score_threshold,
- offset=self.offset,
- doc_start_delimiter="",
- doc_end_delimiter="",
- return_type="both",
- **self.kwargs,
- )
-
- return self._extract_documents(response)
-
- async def _aget_relevant_documents(
- self, query: str, *, run_manager: AsyncCallbackManagerForRetrieverRun
- ) -> List[Document]:
- """Asynchronously get documents relevant to a query.
- Args:
- query: String to find relevant documents for
- run_manager: The callbacks handler to use
- Returns:
- List of relevant documents
- """
- try:
- from asknews_sdk import AsyncAskNewsSDK
- except ImportError:
- raise ImportError(
- "AskNews python package not found. "
- "Please install it with `pip install asknews`."
- )
- an_client = AsyncAskNewsSDK(
- client_id=self.client_id or os.environ["ASKNEWS_CLIENT_ID"],
- client_secret=self.client_secret or os.environ["ASKNEWS_CLIENT_SECRET"],
- scopes=["news"],
- )
- response = await an_client.news.search_news(
- query=query,
- n_articles=self.k,
- start_timestamp=self.start_timestamp,
- end_timestamp=self.end_timestamp,
- method=self.method,
- categories=self.categories,
- historical=self.historical,
- similarity_score_threshold=self.similarity_score_threshold,
- offset=self.offset,
- return_type="both",
- doc_start_delimiter="",
- doc_end_delimiter="",
- **self.kwargs,
- )
-
- return self._extract_documents(response)
-
- def _extract_documents(self, response: Any) -> List[Document]:
- """Extract documents from an api response."""
-
- from asknews_sdk.dto.news import SearchResponse
-
- sr: SearchResponse = response
- matches = re.findall(r"(.*?)", sr.as_string, re.DOTALL)
- docs = [
- Document(
- page_content=matches[i].strip(),
- metadata={
- "title": sr.as_dicts[i].title,
- "source": str(sr.as_dicts[i].article_url)
- if sr.as_dicts[i].article_url
- else None,
- "images": sr.as_dicts[i].image_url,
- },
- )
- for i in range(len(matches))
- ]
- return docs
diff --git a/libs/community/langchain_community/retrievers/azure_ai_search.py b/libs/community/langchain_community/retrievers/azure_ai_search.py
deleted file mode 100644
index 6db86a3c3b..0000000000
--- a/libs/community/langchain_community/retrievers/azure_ai_search.py
+++ /dev/null
@@ -1,226 +0,0 @@
-from __future__ import annotations
-
-import json
-from typing import Any, Dict, List, Optional
-
-import aiohttp
-import requests
-from langchain_core.callbacks import (
- AsyncCallbackManagerForRetrieverRun,
- CallbackManagerForRetrieverRun,
-)
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from langchain_core.utils import get_from_dict_or_env, get_from_env
-from pydantic import ConfigDict, model_validator
-
-DEFAULT_URL_SUFFIX = "search.windows.net"
-"""Default URL Suffix for endpoint connection - commercial cloud"""
-
-
-class AzureAISearchRetriever(BaseRetriever):
- """`Azure AI Search` service retriever.
-
- Setup:
- See here for more detail: https://python.langchain.com/docs/integrations/retrievers/azure_ai_search/
-
- We will need to install the below dependencies and set the required
- environment variables:
-
- .. code-block:: bash
-
- pip install -U langchain-community azure-identity azure-search-documents
- export AZURE_AI_SEARCH_SERVICE_NAME=""
- export AZURE_AI_SEARCH_INDEX_NAME=""
-
- export AZURE_AI_SEARCH_API_KEY=""
- or
- export AZURE_AI_SEARCH_BEARER_TOKEN=""
-
- Key init args:
- content_key: str
- top_k: int
- index_name: str
-
- Instantiate:
- .. code-block:: python
-
- from langchain_community.retrievers import AzureAISearchRetriever
-
- retriever = AzureAISearchRetriever(
- content_key="content", top_k=1, index_name="langchain-vector-demo"
- )
-
- Usage:
- .. code-block:: python
-
- retriever.invoke("here is my unstructured query string")
-
- Use within a chain:
- .. code-block:: python
-
- from langchain_core.output_parsers import StrOutputParser
- from langchain_core.prompts import ChatPromptTemplate
- from langchain_core.runnables import RunnablePassthrough
- from langchain_openai import AzureChatOpenAI
-
- prompt = ChatPromptTemplate.from_template(
- \"\"\"Answer the question based only on the context provided.
-
- Context: {context}
-
- Question: {question}\"\"\"
- )
-
- llm = AzureChatOpenAI(azure_deployment="gpt-35-turbo")
-
- def format_docs(docs):
- return "\\n\\n".join(doc.page_content for doc in docs)
-
- chain = (
- {"context": retriever | format_docs, "question": RunnablePassthrough()}
- | prompt
- | llm
- | StrOutputParser()
- )
-
- chain.invoke("...")
-
- """ # noqa: E501
-
- service_name: str = ""
- """Name of Azure AI Search service"""
- index_name: str = ""
- """Name of Index inside Azure AI Search service"""
- api_key: str = ""
- """API Key. Both Admin and Query keys work, but for reading data it's
- recommended to use a Query key."""
- api_version: str = "2023-11-01"
- """API version"""
- aiosession: Optional[aiohttp.ClientSession] = None
- """ClientSession, in case we want to reuse connection for better performance."""
- azure_ad_token: str = ""
- """Your Azure Active Directory token.
-
- Automatically inferred from env var `AZURE_AI_SEARCH_AD_TOKEN` if not provided.
-
- For more:
- https://www.microsoft.com/en-us/security/business/identity-access/microsoft-entra-id.
- """
- content_key: str = "content"
- """Key in a retrieved result to set as the Document page_content."""
- top_k: Optional[int] = None
- """Number of results to retrieve. Set to None to retrieve all results."""
- filter: Optional[str] = None
- """OData $filter expression to apply to the search query."""
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- extra="forbid",
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that service name, index name and api key exists in environment."""
- values["service_name"] = get_from_dict_or_env(
- values, "service_name", "AZURE_AI_SEARCH_SERVICE_NAME"
- )
- values["index_name"] = get_from_dict_or_env(
- values, "index_name", "AZURE_AI_SEARCH_INDEX_NAME"
- )
- values["azure_ad_token"] = get_from_dict_or_env(
- values, "azure_ad_token", "AZURE_AI_SEARCH_AD_TOKEN", default=""
- )
- values["api_key"] = get_from_dict_or_env(
- values, "api_key", "AZURE_AI_SEARCH_API_KEY", default=""
- )
- if values["azure_ad_token"] == "" and values["api_key"] == "":
- raise ValueError(
- "Missing credentials. Please pass one of `api_key`, `azure_ad_token`, "
- "or the `AZURE_AI_SEARCH_API_KEY` or `AZURE_AI_SEARCH_AD_TOKEN` "
- "environment variables."
- )
-
- return values
-
- def _build_search_url(self, query: str) -> str:
- url_suffix = get_from_env("", "AZURE_AI_SEARCH_URL_SUFFIX", DEFAULT_URL_SUFFIX)
- if url_suffix in self.service_name and "https://" in self.service_name:
- base_url = f"{self.service_name}/"
- elif url_suffix in self.service_name and "https://" not in self.service_name:
- base_url = f"https://{self.service_name}/"
- elif url_suffix not in self.service_name and "https://" in self.service_name:
- base_url = f"{self.service_name}.{url_suffix}/"
- elif (
- url_suffix not in self.service_name and "https://" not in self.service_name
- ):
- base_url = f"https://{self.service_name}.{url_suffix}/"
- else:
- # pass to Azure to throw a specific error
- base_url = self.service_name
- endpoint_path = f"indexes/{self.index_name}/docs?api-version={self.api_version}"
- top_param = f"&$top={self.top_k}" if self.top_k else ""
- filter_param = f"&$filter={self.filter}" if self.filter else ""
- return base_url + endpoint_path + f"&search={query}" + top_param + filter_param
-
- @property
- def _headers(self) -> Dict[str, str]:
- headers = {
- "Content-Type": "application/json",
- }
- if not self.azure_ad_token:
- headers["Authorization"] = f"Bearer {self.azure_ad_token}"
- else:
- headers["api-key"] = f"{self.api_key}"
- return headers
-
- def _search(self, query: str) -> List[dict]:
- search_url = self._build_search_url(query)
- response = requests.get(search_url, headers=self._headers)
- if response.status_code != 200:
- raise Exception(f"Error in search request: {response}")
-
- return json.loads(response.text)["value"]
-
- async def _asearch(self, query: str) -> List[dict]:
- search_url = self._build_search_url(query)
- if not self.aiosession:
- async with aiohttp.ClientSession() as session:
- async with session.get(search_url, headers=self._headers) as response:
- response_json = await response.json()
- else:
- async with self.aiosession.get(
- search_url, headers=self._headers
- ) as response:
- response_json = await response.json()
-
- return response_json["value"]
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- search_results = self._search(query)
-
- return [
- Document(page_content=result.pop(self.content_key), metadata=result)
- for result in search_results
- ]
-
- async def _aget_relevant_documents(
- self, query: str, *, run_manager: AsyncCallbackManagerForRetrieverRun
- ) -> List[Document]:
- search_results = await self._asearch(query)
-
- return [
- Document(page_content=result.pop(self.content_key), metadata=result)
- for result in search_results
- ]
-
-
-# For backwards compatibility
-class AzureCognitiveSearchRetriever(AzureAISearchRetriever):
- """`Azure Cognitive Search` service retriever.
- This version of the retriever will soon be
- depreciated. Please switch to AzureAISearchRetriever
- """
diff --git a/libs/community/langchain_community/retrievers/bedrock.py b/libs/community/langchain_community/retrievers/bedrock.py
deleted file mode 100644
index 3dd33a045d..0000000000
--- a/libs/community/langchain_community/retrievers/bedrock.py
+++ /dev/null
@@ -1,186 +0,0 @@
-from typing import Any, Dict, List, Optional
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from pydantic import BaseModel, model_validator
-
-
-class VectorSearchConfig(BaseModel, extra="allow"):
- """Configuration for vector search."""
-
- numberOfResults: int = 4
-
-
-class RetrievalConfig(BaseModel, extra="allow"):
- """Configuration for retrieval."""
-
- vectorSearchConfiguration: VectorSearchConfig
-
-
-@deprecated(
- since="0.3.16",
- removal="1.0",
- alternative_import="langchain_aws.AmazonKnowledgeBasesRetriever",
-)
-class AmazonKnowledgeBasesRetriever(BaseRetriever):
- """Amazon Bedrock Knowledge Bases retriever.
-
- See https://aws.amazon.com/bedrock/knowledge-bases for more info.
-
- Setup:
- Install ``langchain-aws``:
-
- .. code-block:: bash
-
- pip install -U langchain-aws
-
- Key init args:
- knowledge_base_id: Knowledge Base ID.
- region_name: The aws region e.g., `us-west-2`.
- Fallback to AWS_DEFAULT_REGION env variable or region specified in
- ~/.aws/config.
- credentials_profile_name: The name of the profile in the ~/.aws/credentials
- or ~/.aws/config files, which has either access keys or role information
- specified. If not specified, the default credential profile or, if on an
- EC2 instance, credentials from IMDS will be used.
- client: boto3 client for bedrock agent runtime.
- retrieval_config: Configuration for retrieval.
-
- Instantiate:
- .. code-block:: python
-
- from langchain_community.retrievers import AmazonKnowledgeBasesRetriever
-
- retriever = AmazonKnowledgeBasesRetriever(
- knowledge_base_id="",
- retrieval_config={
- "vectorSearchConfiguration": {
- "numberOfResults": 4
- }
- },
- )
-
- Usage:
- .. code-block:: python
-
- query = "..."
-
- retriever.invoke(query)
-
- Use within a chain:
- .. code-block:: python
-
- from langchain_aws import ChatBedrockConverse
- from langchain_core.output_parsers import StrOutputParser
- from langchain_core.prompts import ChatPromptTemplate
- from langchain_core.runnables import RunnablePassthrough
- from langchain_openai import ChatOpenAI
-
- prompt = ChatPromptTemplate.from_template(
- \"\"\"Answer the question based only on the context provided.
-
- Context: {context}
-
- Question: {question}\"\"\"
- )
-
- llm = ChatBedrockConverse(
- model_id="anthropic.claude-3-5-sonnet-20240620-v1:0"
- )
-
- def format_docs(docs):
- return "\\n\\n".join(doc.page_content for doc in docs)
-
- chain = (
- {"context": retriever | format_docs, "question": RunnablePassthrough()}
- | prompt
- | llm
- | StrOutputParser()
- )
-
- chain.invoke("...")
-
- """ # noqa: E501
-
- knowledge_base_id: str
- region_name: Optional[str] = None
- credentials_profile_name: Optional[str] = None
- endpoint_url: Optional[str] = None
- client: Any
- retrieval_config: RetrievalConfig
-
- @model_validator(mode="before")
- @classmethod
- def create_client(cls, values: Dict[str, Any]) -> Any:
- if values.get("client") is not None:
- return values
-
- try:
- import boto3
- from botocore.client import Config
- from botocore.exceptions import UnknownServiceError
-
- if values.get("credentials_profile_name"):
- session = boto3.Session(profile_name=values["credentials_profile_name"])
- else:
- # use default credentials
- session = boto3.Session()
-
- client_params = {
- "config": Config(
- connect_timeout=120, read_timeout=120, retries={"max_attempts": 0}
- )
- }
- if values.get("region_name"):
- client_params["region_name"] = values["region_name"]
-
- if values.get("endpoint_url"):
- client_params["endpoint_url"] = values["endpoint_url"]
-
- values["client"] = session.client("bedrock-agent-runtime", **client_params)
-
- return values
- except ImportError:
- raise ImportError(
- "Could not import boto3 python package. "
- "Please install it with `pip install boto3`."
- )
- except UnknownServiceError as e:
- raise ImportError(
- "Ensure that you have installed the latest boto3 package "
- "that contains the API for `bedrock-runtime-agent`."
- ) from e
- except Exception as e:
- raise ValueError(
- "Could not load credentials to authenticate with AWS client. "
- "Please check that credentials in the specified "
- "profile name are valid."
- ) from e
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- response = self.client.retrieve(
- retrievalQuery={"text": query.strip()},
- knowledgeBaseId=self.knowledge_base_id,
- retrievalConfiguration=self.retrieval_config.dict(),
- )
- results = response["retrievalResults"]
- documents = []
- for result in results:
- content = result["content"]["text"]
- result.pop("content")
- if "score" not in result:
- result["score"] = 0
- if "metadata" in result:
- result["source_metadata"] = result.pop("metadata")
- documents.append(
- Document(
- page_content=content,
- metadata=result,
- )
- )
-
- return documents
diff --git a/libs/community/langchain_community/retrievers/bm25.py b/libs/community/langchain_community/retrievers/bm25.py
deleted file mode 100644
index 70910ce170..0000000000
--- a/libs/community/langchain_community/retrievers/bm25.py
+++ /dev/null
@@ -1,116 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Callable, Dict, Iterable, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from pydantic import ConfigDict, Field
-
-
-def default_preprocessing_func(text: str) -> List[str]:
- return text.split()
-
-
-class BM25Retriever(BaseRetriever):
- """`BM25` retriever without Elasticsearch."""
-
- vectorizer: Any = None
- """ BM25 vectorizer."""
- docs: List[Document] = Field(repr=False)
- """ List of documents."""
- k: int = 4
- """ Number of documents to return."""
- preprocess_func: Callable[[str], List[str]] = default_preprocessing_func
- """ Preprocessing function to use on the text before BM25 vectorization."""
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- )
-
- @classmethod
- def from_texts(
- cls,
- texts: Iterable[str],
- metadatas: Optional[Iterable[dict]] = None,
- ids: Optional[Iterable[str]] = None,
- bm25_params: Optional[Dict[str, Any]] = None,
- preprocess_func: Callable[[str], List[str]] = default_preprocessing_func,
- **kwargs: Any,
- ) -> BM25Retriever:
- """
- Create a BM25Retriever from a list of texts.
- Args:
- texts: A list of texts to vectorize.
- metadatas: A list of metadata dicts to associate with each text.
- ids: A list of ids to associate with each text.
- bm25_params: Parameters to pass to the BM25 vectorizer.
- preprocess_func: A function to preprocess each text before vectorization.
- **kwargs: Any other arguments to pass to the retriever.
-
- Returns:
- A BM25Retriever instance.
- """
- try:
- from rank_bm25 import BM25Okapi
- except ImportError:
- raise ImportError(
- "Could not import rank_bm25, please install with `pip install "
- "rank_bm25`."
- )
-
- texts_processed = [preprocess_func(t) for t in texts]
- bm25_params = bm25_params or {}
- vectorizer = BM25Okapi(texts_processed, **bm25_params)
- metadatas = metadatas or ({} for _ in texts)
- if ids:
- docs = [
- Document(page_content=t, metadata=m, id=i)
- for t, m, i in zip(texts, metadatas, ids)
- ]
- else:
- docs = [
- Document(page_content=t, metadata=m) for t, m in zip(texts, metadatas)
- ]
- return cls(
- vectorizer=vectorizer, docs=docs, preprocess_func=preprocess_func, **kwargs
- )
-
- @classmethod
- def from_documents(
- cls,
- documents: Iterable[Document],
- *,
- bm25_params: Optional[Dict[str, Any]] = None,
- preprocess_func: Callable[[str], List[str]] = default_preprocessing_func,
- **kwargs: Any,
- ) -> BM25Retriever:
- """
- Create a BM25Retriever from a list of Documents.
- Args:
- documents: A list of Documents to vectorize.
- bm25_params: Parameters to pass to the BM25 vectorizer.
- preprocess_func: A function to preprocess each text before vectorization.
- **kwargs: Any other arguments to pass to the retriever.
-
- Returns:
- A BM25Retriever instance.
- """
- texts, metadatas, ids = zip(
- *((d.page_content, d.metadata, d.id) for d in documents)
- )
- return cls.from_texts(
- texts=texts,
- bm25_params=bm25_params,
- metadatas=metadatas,
- ids=ids,
- preprocess_func=preprocess_func,
- **kwargs,
- )
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- processed_query = self.preprocess_func(query)
- return_docs = self.vectorizer.get_top_n(processed_query, self.docs, n=self.k)
- return return_docs
diff --git a/libs/community/langchain_community/retrievers/breebs.py b/libs/community/langchain_community/retrievers/breebs.py
deleted file mode 100644
index b6551b090b..0000000000
--- a/libs/community/langchain_community/retrievers/breebs.py
+++ /dev/null
@@ -1,49 +0,0 @@
-from typing import List
-
-import requests
-from langchain_core.callbacks.manager import CallbackManagerForRetrieverRun
-from langchain_core.documents.base import Document
-from langchain_core.retrievers import BaseRetriever
-
-
-class BreebsRetriever(BaseRetriever):
- """A retriever class for `Breebs`.
-
- See https://www.breebs.com/ for more info.
- Args:
- breeb_key: The key to trigger the breeb
- (specialized knowledge pill on a specific topic).
-
- To retrieve the list of all available Breebs : you can call https://breebs.promptbreeders.com/web/listbreebs
- """
-
- breeb_key: str
- url: str = "https://breebs.promptbreeders.com/knowledge"
-
- def __init__(self, breeb_key: str):
- super().__init__(breeb_key=breeb_key) # type: ignore[call-arg]
- self.breeb_key = breeb_key
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- """Retrieve context for given query.
- Note that for time being there is no score."""
- r = requests.post(
- self.url,
- json={
- "breeb_key": self.breeb_key,
- "query": query,
- },
- )
- if r.status_code != 200:
- return []
- else:
- chunks = r.json()
- return [
- Document(
- page_content=chunk["content"],
- metadata={"source": chunk["source_url"], "score": 1},
- )
- for chunk in chunks
- ]
diff --git a/libs/community/langchain_community/retrievers/chaindesk.py b/libs/community/langchain_community/retrievers/chaindesk.py
deleted file mode 100644
index 4c8aa2c582..0000000000
--- a/libs/community/langchain_community/retrievers/chaindesk.py
+++ /dev/null
@@ -1,92 +0,0 @@
-from typing import Any, List, Optional
-
-import aiohttp
-import requests
-from langchain_core.callbacks import (
- AsyncCallbackManagerForRetrieverRun,
- CallbackManagerForRetrieverRun,
-)
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-
-class ChaindeskRetriever(BaseRetriever):
- """`Chaindesk API` retriever."""
-
- datastore_url: str
- top_k: Optional[int]
- api_key: Optional[str]
-
- def __init__(
- self,
- datastore_url: str,
- top_k: Optional[int] = None,
- api_key: Optional[str] = None,
- ):
- self.datastore_url = datastore_url
- self.api_key = api_key
- self.top_k = top_k
-
- def _get_relevant_documents(
- self,
- query: str,
- *,
- run_manager: CallbackManagerForRetrieverRun,
- **kwargs: Any,
- ) -> List[Document]:
- response = requests.post(
- self.datastore_url,
- json={
- "query": query,
- **({"topK": self.top_k} if self.top_k is not None else {}),
- },
- headers={
- "Content-Type": "application/json",
- **(
- {"Authorization": f"Bearer {self.api_key}"}
- if self.api_key is not None
- else {}
- ),
- },
- )
- data = response.json()
- return [
- Document(
- page_content=r["text"],
- metadata={"source": r["source"], "score": r["score"]},
- )
- for r in data["results"]
- ]
-
- async def _aget_relevant_documents(
- self,
- query: str,
- *,
- run_manager: AsyncCallbackManagerForRetrieverRun,
- **kwargs: Any,
- ) -> List[Document]:
- async with aiohttp.ClientSession() as session:
- async with session.request(
- "POST",
- self.datastore_url,
- json={
- "query": query,
- **({"topK": self.top_k} if self.top_k is not None else {}),
- },
- headers={
- "Content-Type": "application/json",
- **(
- {"Authorization": f"Bearer {self.api_key}"}
- if self.api_key is not None
- else {}
- ),
- },
- ) as response:
- data = await response.json()
- return [
- Document(
- page_content=r["text"],
- metadata={"source": r["source"], "score": r["score"]},
- )
- for r in data["results"]
- ]
diff --git a/libs/community/langchain_community/retrievers/chatgpt_plugin_retriever.py b/libs/community/langchain_community/retrievers/chatgpt_plugin_retriever.py
deleted file mode 100644
index 08559110bc..0000000000
--- a/libs/community/langchain_community/retrievers/chatgpt_plugin_retriever.py
+++ /dev/null
@@ -1,89 +0,0 @@
-from __future__ import annotations
-
-from typing import List, Optional
-
-import aiohttp
-import requests
-from langchain_core.callbacks import (
- AsyncCallbackManagerForRetrieverRun,
- CallbackManagerForRetrieverRun,
-)
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from pydantic import ConfigDict
-
-
-class ChatGPTPluginRetriever(BaseRetriever):
- """`ChatGPT plugin` retriever."""
-
- url: str
- """URL of the ChatGPT plugin."""
- bearer_token: str
- """Bearer token for the ChatGPT plugin."""
- top_k: int = 3
- """Number of documents to return."""
- filter: Optional[dict] = None
- """Filter to apply to the results."""
- aiosession: Optional[aiohttp.ClientSession] = None
- """Aiohttp session to use for requests."""
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- )
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- url, json, headers = self._create_request(query)
- response = requests.post(url, json=json, headers=headers)
- results = response.json()["results"][0]["results"]
- docs = []
- for d in results:
- content = d.pop("text")
- metadata = d.pop("metadata", d)
- if metadata.get("source_id"):
- metadata["source"] = metadata.pop("source_id")
- docs.append(Document(page_content=content, metadata=metadata))
- return docs
-
- async def _aget_relevant_documents(
- self, query: str, *, run_manager: AsyncCallbackManagerForRetrieverRun
- ) -> List[Document]:
- url, json, headers = self._create_request(query)
-
- if not self.aiosession:
- async with aiohttp.ClientSession() as session:
- async with session.post(url, headers=headers, json=json) as response:
- res = await response.json()
- else:
- async with self.aiosession.post(
- url, headers=headers, json=json
- ) as response:
- res = await response.json()
-
- results = res["results"][0]["results"]
- docs = []
- for d in results:
- content = d.pop("text")
- metadata = d.pop("metadata", d)
- if metadata.get("source_id"):
- metadata["source"] = metadata.pop("source_id")
- docs.append(Document(page_content=content, metadata=metadata))
- return docs
-
- def _create_request(self, query: str) -> tuple[str, dict, dict]:
- url = f"{self.url}/query"
- json = {
- "queries": [
- {
- "query": query,
- "filter": self.filter,
- "top_k": self.top_k,
- }
- ]
- }
- headers = {
- "Content-Type": "application/json",
- "Authorization": f"Bearer {self.bearer_token}",
- }
- return url, json, headers
diff --git a/libs/community/langchain_community/retrievers/cohere_rag_retriever.py b/libs/community/langchain_community/retrievers/cohere_rag_retriever.py
deleted file mode 100644
index f76aafa290..0000000000
--- a/libs/community/langchain_community/retrievers/cohere_rag_retriever.py
+++ /dev/null
@@ -1,96 +0,0 @@
-from __future__ import annotations
-
-from typing import TYPE_CHECKING, Any, Dict, List
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import (
- AsyncCallbackManagerForRetrieverRun,
- CallbackManagerForRetrieverRun,
-)
-from langchain_core.documents import Document
-from langchain_core.language_models.chat_models import BaseChatModel
-from langchain_core.messages import HumanMessage
-from langchain_core.retrievers import BaseRetriever
-from pydantic import ConfigDict, Field
-
-if TYPE_CHECKING:
- from langchain_core.messages import BaseMessage
-
-
-def _get_docs(response: Any) -> List[Document]:
- docs = (
- []
- if "documents" not in response.generation_info
- else [
- Document(page_content=doc["snippet"], metadata=doc)
- for doc in response.generation_info["documents"]
- ]
- )
- docs.append(
- Document(
- page_content=response.message.content,
- metadata={
- "type": "model_response",
- "citations": response.generation_info["citations"],
- "search_results": response.generation_info["search_results"],
- "search_queries": response.generation_info["search_queries"],
- "token_count": response.generation_info["token_count"],
- },
- )
- )
- return docs
-
-
-@deprecated(
- since="0.0.30",
- removal="1.0",
- alternative_import="langchain_cohere.CohereRagRetriever",
-)
-class CohereRagRetriever(BaseRetriever):
- """Cohere Chat API with RAG."""
-
- connectors: List[Dict] = Field(default_factory=lambda: [{"id": "web-search"}])
- """
- When specified, the model's reply will be enriched with information found by
- querying each of the connectors (RAG). These will be returned as langchain
- documents.
-
- Currently only accepts {"id": "web-search"}.
- """
-
- llm: BaseChatModel
- """Cohere ChatModel to use."""
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- )
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun, **kwargs: Any
- ) -> List[Document]:
- messages: List[List[BaseMessage]] = [[HumanMessage(content=query)]]
- res = self.llm.generate(
- messages,
- connectors=self.connectors,
- callbacks=run_manager.get_child(),
- **kwargs,
- ).generations[0][0]
- return _get_docs(res)
-
- async def _aget_relevant_documents(
- self,
- query: str,
- *,
- run_manager: AsyncCallbackManagerForRetrieverRun,
- **kwargs: Any,
- ) -> List[Document]:
- messages: List[List[BaseMessage]] = [[HumanMessage(content=query)]]
- res = (
- await self.llm.agenerate(
- messages,
- connectors=self.connectors,
- callbacks=run_manager.get_child(),
- **kwargs,
- )
- ).generations[0][0]
- return _get_docs(res)
diff --git a/libs/community/langchain_community/retrievers/databerry.py b/libs/community/langchain_community/retrievers/databerry.py
deleted file mode 100644
index c1ea627700..0000000000
--- a/libs/community/langchain_community/retrievers/databerry.py
+++ /dev/null
@@ -1,74 +0,0 @@
-from typing import List, Optional
-
-import aiohttp
-import requests
-from langchain_core.callbacks import (
- AsyncCallbackManagerForRetrieverRun,
- CallbackManagerForRetrieverRun,
-)
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-
-class DataberryRetriever(BaseRetriever):
- """`Databerry API` retriever."""
-
- datastore_url: str
- top_k: Optional[int]
- api_key: Optional[str]
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- response = requests.post(
- self.datastore_url,
- json={
- "query": query,
- **({"topK": self.top_k} if self.top_k is not None else {}),
- },
- headers={
- "Content-Type": "application/json",
- **(
- {"Authorization": f"Bearer {self.api_key}"}
- if self.api_key is not None
- else {}
- ),
- },
- )
- data = response.json()
- return [
- Document(
- page_content=r["text"],
- metadata={"source": r["source"], "score": r["score"]},
- )
- for r in data["results"]
- ]
-
- async def _aget_relevant_documents(
- self, query: str, *, run_manager: AsyncCallbackManagerForRetrieverRun
- ) -> List[Document]:
- async with aiohttp.ClientSession() as session:
- async with session.request(
- "POST",
- self.datastore_url,
- json={
- "query": query,
- **({"topK": self.top_k} if self.top_k is not None else {}),
- },
- headers={
- "Content-Type": "application/json",
- **(
- {"Authorization": f"Bearer {self.api_key}"}
- if self.api_key is not None
- else {}
- ),
- },
- ) as response:
- data = await response.json()
- return [
- Document(
- page_content=r["text"],
- metadata={"source": r["source"], "score": r["score"]},
- )
- for r in data["results"]
- ]
diff --git a/libs/community/langchain_community/retrievers/docarray.py b/libs/community/langchain_community/retrievers/docarray.py
deleted file mode 100644
index 2e602c1065..0000000000
--- a/libs/community/langchain_community/retrievers/docarray.py
+++ /dev/null
@@ -1,208 +0,0 @@
-from enum import Enum
-from typing import Any, Dict, List, Optional, Union
-
-import numpy as np
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.embeddings import Embeddings
-from langchain_core.retrievers import BaseRetriever
-from langchain_core.utils.pydantic import get_fields
-from pydantic import ConfigDict
-
-from langchain_community.vectorstores.utils import maximal_marginal_relevance
-
-
-class SearchType(str, Enum):
- """Enumerator of the types of search to perform."""
-
- similarity = "similarity"
- mmr = "mmr"
-
-
-class DocArrayRetriever(BaseRetriever):
- """`DocArray Document Indices` retriever.
-
- Currently, it supports 5 backends:
- InMemoryExactNNIndex, HnswDocumentIndex, QdrantDocumentIndex,
- ElasticDocIndex, and WeaviateDocumentIndex.
-
- Args:
- index: One of the above-mentioned index instances
- embeddings: Embedding model to represent text as vectors
- search_field: Field to consider for searching in the documents.
- Should be an embedding/vector/tensor.
- content_field: Field that represents the main content in your document schema.
- Will be used as a `page_content`. Everything else will go into `metadata`.
- search_type: Type of search to perform (similarity / mmr)
- filters: Filters applied for document retrieval.
- top_k: Number of documents to return
- """
-
- index: Any = None
- embeddings: Embeddings
- search_field: str
- content_field: str
- search_type: SearchType = SearchType.similarity
- top_k: int = 1
- filters: Optional[Any] = None
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- )
-
- def _get_relevant_documents(
- self,
- query: str,
- *,
- run_manager: CallbackManagerForRetrieverRun,
- ) -> List[Document]:
- """Get documents relevant for a query.
-
- Args:
- query: string to find relevant documents for
-
- Returns:
- List of relevant documents
- """
- query_emb = np.array(self.embeddings.embed_query(query))
-
- if self.search_type == SearchType.similarity:
- results = self._similarity_search(query_emb)
- elif self.search_type == SearchType.mmr:
- results = self._mmr_search(query_emb)
- else:
- raise ValueError(
- f"Search type {self.search_type} does not exist. "
- f"Choose either 'similarity' or 'mmr'."
- )
-
- return results
-
- def _search(
- self, query_emb: np.ndarray, top_k: int
- ) -> List[Union[Dict[str, Any], Any]]:
- """
- Perform a search using the query embedding and return top_k documents.
-
- Args:
- query_emb: Query represented as an embedding
- top_k: Number of documents to return
-
- Returns:
- A list of top_k documents matching the query
- """
-
- from docarray.index import ElasticDocIndex, WeaviateDocumentIndex
-
- filter_args = {}
- search_field = self.search_field
- if isinstance(self.index, WeaviateDocumentIndex):
- filter_args["where_filter"] = self.filters
- search_field = ""
- elif isinstance(self.index, ElasticDocIndex):
- filter_args["query"] = self.filters
- else:
- filter_args["filter_query"] = self.filters
-
- if self.filters:
- query = (
- self.index.build_query() # get empty query object
- .find(
- query=query_emb, search_field=search_field
- ) # add vector similarity search
- .filter(**filter_args) # add filter search
- .build(limit=top_k) # build the query
- )
- # execute the combined query and return the results
- docs = self.index.execute_query(query)
- if hasattr(docs, "documents"):
- docs = docs.documents
- docs = docs[:top_k]
- else:
- docs = self.index.find(
- query=query_emb, search_field=search_field, limit=top_k
- ).documents
- return docs
-
- def _similarity_search(self, query_emb: np.ndarray) -> List[Document]:
- """
- Perform a similarity search.
-
- Args:
- query_emb: Query represented as an embedding
-
- Returns:
- A list of documents most similar to the query
- """
- docs = self._search(query_emb=query_emb, top_k=self.top_k)
- results = [self._docarray_to_langchain_doc(doc) for doc in docs]
- return results
-
- def _mmr_search(self, query_emb: np.ndarray) -> List[Document]:
- """
- Perform a maximal marginal relevance (mmr) search.
-
- Args:
- query_emb: Query represented as an embedding
-
- Returns:
- A list of diverse documents related to the query
- """
- docs = self._search(query_emb=query_emb, top_k=20)
-
- mmr_selected = maximal_marginal_relevance(
- query_emb,
- [
- doc[self.search_field]
- if isinstance(doc, dict)
- else getattr(doc, self.search_field)
- for doc in docs
- ],
- k=self.top_k,
- )
- results = [self._docarray_to_langchain_doc(docs[idx]) for idx in mmr_selected]
- return results
-
- def _docarray_to_langchain_doc(self, doc: Union[Dict[str, Any], Any]) -> Document:
- """
- Convert a DocArray document (which also might be a dict)
- to a langchain document format.
-
- DocArray document can contain arbitrary fields, so the mapping is done
- in the following way:
-
- page_content <-> content_field
- metadata <-> all other fields excluding
- tensors and embeddings (so float, int, string)
-
- Args:
- doc: DocArray document
-
- Returns:
- Document in langchain format
-
- Raises:
- ValueError: If the document doesn't contain the content field
- """
-
- fields = doc.keys() if isinstance(doc, dict) else get_fields(doc)
-
- if self.content_field not in fields:
- raise ValueError(
- f"Document does not contain the content field - {self.content_field}."
- )
- lc_doc = Document(
- page_content=doc[self.content_field]
- if isinstance(doc, dict)
- else getattr(doc, self.content_field)
- )
-
- for name in fields:
- value = doc[name] if isinstance(doc, dict) else getattr(doc, name)
- if (
- isinstance(value, (str, int, float, bool))
- and name != self.content_field
- ):
- lc_doc.metadata[name] = value
-
- return lc_doc
diff --git a/libs/community/langchain_community/retrievers/dria_index.py b/libs/community/langchain_community/retrievers/dria_index.py
deleted file mode 100644
index 8f3e287d8e..0000000000
--- a/libs/community/langchain_community/retrievers/dria_index.py
+++ /dev/null
@@ -1,87 +0,0 @@
-"""Wrapper around Dria Retriever."""
-
-from typing import Any, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-from langchain_community.utilities import DriaAPIWrapper
-
-
-class DriaRetriever(BaseRetriever):
- """`Dria` retriever using the DriaAPIWrapper."""
-
- api_wrapper: DriaAPIWrapper
-
- def __init__(self, api_key: str, contract_id: Optional[str] = None, **kwargs: Any):
- """
- Initialize the DriaRetriever with a DriaAPIWrapper instance.
-
- Args:
- api_key: The API key for Dria.
- contract_id: The contract ID of the knowledge base to interact with.
- """
- api_wrapper = DriaAPIWrapper(api_key=api_key, contract_id=contract_id)
- super().__init__(api_wrapper=api_wrapper, **kwargs) # type: ignore[call-arg]
-
- def create_knowledge_base(
- self,
- name: str,
- description: str,
- category: str = "Unspecified",
- embedding: str = "jina",
- ) -> str:
- """Create a new knowledge base in Dria.
-
- Args:
- name: The name of the knowledge base.
- description: The description of the knowledge base.
- category: The category of the knowledge base.
- embedding: The embedding model to use for the knowledge base.
-
-
- Returns:
- The ID of the created knowledge base.
- """
- response = self.api_wrapper.create_knowledge_base(
- name, description, category, embedding
- )
- return response
-
- def add_texts(
- self,
- texts: List,
- ) -> None:
- """Add texts to the Dria knowledge base.
-
- Args:
- texts: An iterable of texts and metadatas to add to the knowledge base.
-
- Returns:
- List of IDs representing the added texts.
- """
- data = [{"text": text["text"], "metadata": text["metadata"]} for text in texts]
- self.api_wrapper.insert_data(data)
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- """Retrieve relevant documents from Dria based on a query.
-
- Args:
- query: The query string to search for in the knowledge base.
- run_manager: Callback manager for the retriever run.
-
- Returns:
- A list of Documents containing the search results.
- """
- results = self.api_wrapper.search(query)
- docs = [
- Document(
- page_content=result["metadata"],
- metadata={"id": result["id"], "score": result["score"]},
- )
- for result in results
- ]
- return docs
diff --git a/libs/community/langchain_community/retrievers/elastic_search_bm25.py b/libs/community/langchain_community/retrievers/elastic_search_bm25.py
deleted file mode 100644
index a95264df10..0000000000
--- a/libs/community/langchain_community/retrievers/elastic_search_bm25.py
+++ /dev/null
@@ -1,137 +0,0 @@
-"""Wrapper around Elasticsearch vector database."""
-
-from __future__ import annotations
-
-import uuid
-from typing import Any, Iterable, List
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-
-class ElasticSearchBM25Retriever(BaseRetriever):
- """`Elasticsearch` retriever that uses `BM25`.
-
- To connect to an Elasticsearch instance that requires login credentials,
- including Elastic Cloud, use the Elasticsearch URL format
- https://username:password@es_host:9243. For example, to connect to Elastic
- Cloud, create the Elasticsearch URL with the required authentication details and
- pass it to the ElasticVectorSearch constructor as the named parameter
- elasticsearch_url.
-
- You can obtain your Elastic Cloud URL and login credentials by logging in to the
- Elastic Cloud console at https://cloud.elastic.co, selecting your deployment, and
- navigating to the "Deployments" page.
-
- To obtain your Elastic Cloud password for the default "elastic" user:
-
- 1. Log in to the Elastic Cloud console at https://cloud.elastic.co
- 2. Go to "Security" > "Users"
- 3. Locate the "elastic" user and click "Edit"
- 4. Click "Reset password"
- 5. Follow the prompts to reset the password
-
- The format for Elastic Cloud URLs is
- https://username:password@cluster_id.region_id.gcp.cloud.es.io:9243.
- """
-
- client: Any
- """Elasticsearch client."""
- index_name: str
- """Name of the index to use in Elasticsearch."""
-
- @classmethod
- def create(
- cls, elasticsearch_url: str, index_name: str, k1: float = 2.0, b: float = 0.75
- ) -> ElasticSearchBM25Retriever:
- """
- Create a ElasticSearchBM25Retriever from a list of texts.
-
- Args:
- elasticsearch_url: URL of the Elasticsearch instance to connect to.
- index_name: Name of the index to use in Elasticsearch.
- k1: BM25 parameter k1.
- b: BM25 parameter b.
-
- Returns:
-
- """
- from elasticsearch import Elasticsearch
-
- # Create an Elasticsearch client instance
- es = Elasticsearch(elasticsearch_url)
-
- # Define the index settings and mappings
- settings = {
- "analysis": {"analyzer": {"default": {"type": "standard"}}},
- "similarity": {
- "custom_bm25": {
- "type": "BM25",
- "k1": k1,
- "b": b,
- }
- },
- }
- mappings = {
- "properties": {
- "content": {
- "type": "text",
- "similarity": "custom_bm25", # Use the custom BM25 similarity
- }
- }
- }
-
- # Create the index with the specified settings and mappings
- es.indices.create(index=index_name, mappings=mappings, settings=settings)
- return cls(client=es, index_name=index_name)
-
- def add_texts(
- self,
- texts: Iterable[str],
- refresh_indices: bool = True,
- ) -> List[str]:
- """Run more texts through the embeddings and add to the retriever.
-
- Args:
- texts: Iterable of strings to add to the retriever.
- refresh_indices: bool to refresh ElasticSearch indices
-
- Returns:
- List of ids from adding the texts into the retriever.
- """
- try:
- from elasticsearch.helpers import bulk
- except ImportError:
- raise ImportError(
- "Could not import elasticsearch python package. "
- "Please install it with `pip install elasticsearch`."
- )
- requests = []
- ids = []
- for i, text in enumerate(texts):
- _id = str(uuid.uuid4())
- request = {
- "_op_type": "index",
- "_index": self.index_name,
- "content": text,
- "_id": _id,
- }
- ids.append(_id)
- requests.append(request)
- bulk(self.client, requests)
-
- if refresh_indices:
- self.client.indices.refresh(index=self.index_name)
- return ids
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- query_dict = {"query": {"match": {"content": query}}}
- res = self.client.search(index=self.index_name, body=query_dict)
-
- docs = []
- for r in res["hits"]["hits"]:
- docs.append(Document(page_content=r["_source"]["content"]))
- return docs
diff --git a/libs/community/langchain_community/retrievers/embedchain.py b/libs/community/langchain_community/retrievers/embedchain.py
deleted file mode 100644
index 9c64f628e1..0000000000
--- a/libs/community/langchain_community/retrievers/embedchain.py
+++ /dev/null
@@ -1,74 +0,0 @@
-"""Wrapper around Embedchain Retriever."""
-
-from __future__ import annotations
-
-from typing import Any, Iterable, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-
-class EmbedchainRetriever(BaseRetriever):
- """`Embedchain` retriever."""
-
- client: Any
- """Embedchain Pipeline."""
-
- @classmethod
- def create(cls, yaml_path: Optional[str] = None) -> EmbedchainRetriever:
- """
- Create a EmbedchainRetriever from a YAML configuration file.
-
- Args:
- yaml_path: Path to the YAML configuration file. If not provided,
- a default configuration is used.
-
- Returns:
- An instance of EmbedchainRetriever.
-
- """
- from embedchain import Pipeline
-
- # Create an Embedchain Pipeline instance
- if yaml_path:
- client = Pipeline.from_config(yaml_path=yaml_path)
- else:
- client = Pipeline()
- return cls(client=client)
-
- def add_texts(
- self,
- texts: Iterable[str],
- ) -> List[str]:
- """Run more texts through the embeddings and add to the retriever.
-
- Args:
- texts: Iterable of strings/URLs to add to the retriever.
-
- Returns:
- List of ids from adding the texts into the retriever.
- """
- ids = []
- for text in texts:
- _id = self.client.add(text)
- ids.append(_id)
- return ids
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- res = self.client.search(query)
-
- docs = []
- for r in res:
- docs.append(
- Document(
- page_content=r["context"],
- metadata={
- "source": r["metadata"]["url"],
- "document_id": r["metadata"]["doc_id"],
- },
- )
- )
- return docs
diff --git a/libs/community/langchain_community/retrievers/google_cloud_documentai_warehouse.py b/libs/community/langchain_community/retrievers/google_cloud_documentai_warehouse.py
deleted file mode 100644
index 869602229e..0000000000
--- a/libs/community/langchain_community/retrievers/google_cloud_documentai_warehouse.py
+++ /dev/null
@@ -1,126 +0,0 @@
-"""Retriever wrapper for Google Cloud Document AI Warehouse."""
-
-from typing import TYPE_CHECKING, Any, Dict, List, Optional
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from langchain_core.utils import get_from_dict_or_env, pre_init
-
-from langchain_community.utilities.vertexai import get_client_info
-
-if TYPE_CHECKING:
- from google.cloud.contentwarehouse_v1 import (
- DocumentServiceClient,
- RequestMetadata,
- SearchDocumentsRequest,
- )
- from google.cloud.contentwarehouse_v1.services.document_service.pagers import (
- SearchDocumentsPager,
- )
-
-
-@deprecated(
- since="0.0.32",
- removal="1.0",
- alternative_import="langchain_google_community.DocumentAIWarehouseRetriever",
-)
-class GoogleDocumentAIWarehouseRetriever(BaseRetriever):
- """A retriever based on Document AI Warehouse.
-
- Documents should be created and documents should be uploaded
- in a separate flow, and this retriever uses only Document AI
- schema_id provided to search for relevant documents.
-
- More info: https://cloud.google.com/document-ai-warehouse.
- """
-
- location: str = "us"
- """Google Cloud location where Document AI Warehouse is placed."""
- project_number: str
- """Google Cloud project number, should contain digits only."""
- schema_id: Optional[str] = None
- """Document AI Warehouse schema to query against.
- If nothing is provided, all documents in the project will be searched."""
- qa_size_limit: int = 5
- """The limit on the number of documents returned."""
- client: "DocumentServiceClient" = None #: :meta private:
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validates the environment."""
- try:
- from google.cloud.contentwarehouse_v1 import DocumentServiceClient
- except ImportError as exc:
- raise ImportError(
- "google.cloud.contentwarehouse is not installed."
- "Please install it with pip install google-cloud-contentwarehouse"
- ) from exc
-
- values["project_number"] = get_from_dict_or_env(
- values, "project_number", "PROJECT_NUMBER"
- )
- values["client"] = DocumentServiceClient(
- client_info=get_client_info(module="document-ai-warehouse")
- )
- return values
-
- def _prepare_request_metadata(self, user_ldap: str) -> "RequestMetadata":
- from google.cloud.contentwarehouse_v1 import RequestMetadata, UserInfo
-
- user_info = UserInfo(id=f"user:{user_ldap}")
- return RequestMetadata(user_info=user_info)
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun, **kwargs: Any
- ) -> List[Document]:
- request = self._prepare_search_request(query, **kwargs)
- response = self.client.search_documents(request=request)
- return self._parse_search_response(response=response)
-
- def _prepare_search_request(
- self, query: str, **kwargs: Any
- ) -> "SearchDocumentsRequest":
- from google.cloud.contentwarehouse_v1 import (
- DocumentQuery,
- SearchDocumentsRequest,
- )
-
- try:
- user_ldap = kwargs["user_ldap"]
- except KeyError:
- raise ValueError("Argument user_ldap should be provided!")
-
- request_metadata = self._prepare_request_metadata(user_ldap=user_ldap)
- schemas = []
- if self.schema_id:
- schemas.append(
- self.client.document_schema_path(
- project=self.project_number,
- location=self.location,
- document_schema=self.schema_id,
- )
- )
- return SearchDocumentsRequest(
- parent=self.client.common_location_path(self.project_number, self.location),
- request_metadata=request_metadata,
- document_query=DocumentQuery(
- query=query, is_nl_query=True, document_schema_names=schemas
- ),
- qa_size_limit=self.qa_size_limit,
- )
-
- def _parse_search_response(
- self, response: "SearchDocumentsPager"
- ) -> List[Document]:
- documents = []
- for doc in response.matching_documents:
- metadata = {
- "title": doc.document.title,
- "source": doc.document.raw_document_path,
- }
- documents.append(
- Document(page_content=doc.search_text_snippet, metadata=metadata)
- )
- return documents
diff --git a/libs/community/langchain_community/retrievers/google_vertex_ai_search.py b/libs/community/langchain_community/retrievers/google_vertex_ai_search.py
deleted file mode 100644
index 1cc261695d..0000000000
--- a/libs/community/langchain_community/retrievers/google_vertex_ai_search.py
+++ /dev/null
@@ -1,491 +0,0 @@
-"""Retriever wrapper for Google Vertex AI Search."""
-
-from __future__ import annotations
-
-from typing import TYPE_CHECKING, Any, Dict, List, Optional, Sequence, Tuple
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import BaseModel, ConfigDict, Field, model_validator
-
-from langchain_community.utilities.vertexai import get_client_info
-
-if TYPE_CHECKING:
- from google.api_core.client_options import ClientOptions
- from google.cloud.discoveryengine_v1beta import SearchRequest, SearchResult
-
-
-class _BaseGoogleVertexAISearchRetriever(BaseModel):
- project_id: str
- """Google Cloud Project ID."""
- data_store_id: Optional[str] = None
- """Vertex AI Search data store ID."""
- search_engine_id: Optional[str] = None
- """Vertex AI Search app ID."""
- location_id: str = "global"
- """Vertex AI Search data store location."""
- serving_config_id: str = "default_config"
- """Vertex AI Search serving config ID."""
- credentials: Any = None
- """The default custom credentials (google.auth.credentials.Credentials) to use
- when making API calls. If not provided, credentials will be ascertained from
- the environment."""
- engine_data_type: int = Field(default=0, ge=0, le=3)
- """ Defines the Vertex AI Search app data type
- 0 - Unstructured data
- 1 - Structured data
- 2 - Website data
- 3 - Blended search
- """
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validates the environment."""
- try:
- from google.cloud import discoveryengine_v1beta # noqa: F401
- except ImportError as exc:
- raise ImportError(
- "google.cloud.discoveryengine is not installed."
- "Please install it with pip install "
- "google-cloud-discoveryengine>=0.11.10"
- ) from exc
- try:
- from google.api_core.exceptions import InvalidArgument # noqa: F401
- except ImportError as exc:
- raise ImportError(
- "google.api_core.exceptions is not installed. "
- "Please install it with pip install google-api-core"
- ) from exc
-
- values["project_id"] = get_from_dict_or_env(values, "project_id", "PROJECT_ID")
-
- try:
- values["data_store_id"] = get_from_dict_or_env(
- values, "data_store_id", "DATA_STORE_ID"
- )
- values["search_engine_id"] = get_from_dict_or_env(
- values, "search_engine_id", "SEARCH_ENGINE_ID"
- )
- except Exception:
- pass
-
- return values
-
- @property
- def client_options(self) -> "ClientOptions":
- from google.api_core.client_options import ClientOptions
-
- return ClientOptions(
- api_endpoint=(
- f"{self.location_id}-discoveryengine.googleapis.com"
- if self.location_id != "global"
- else None
- )
- )
-
- def _convert_structured_search_response(
- self, results: Sequence[SearchResult]
- ) -> List[Document]:
- """Converts a sequence of search results to a list of LangChain documents."""
- import json
-
- from google.protobuf.json_format import MessageToDict
-
- documents: List[Document] = []
-
- for result in results:
- document_dict = MessageToDict(
- result.document._pb, preserving_proto_field_name=True
- )
-
- documents.append(
- Document(
- page_content=json.dumps(document_dict.get("struct_data", {})),
- metadata={"id": document_dict["id"], "name": document_dict["name"]},
- )
- )
-
- return documents
-
- def _convert_unstructured_search_response(
- self, results: Sequence[SearchResult], chunk_type: str
- ) -> List[Document]:
- """Converts a sequence of search results to a list of LangChain documents."""
- from google.protobuf.json_format import MessageToDict
-
- documents: List[Document] = []
-
- for result in results:
- document_dict = MessageToDict(
- result.document._pb, preserving_proto_field_name=True
- )
- derived_struct_data = document_dict.get("derived_struct_data")
- if not derived_struct_data:
- continue
-
- doc_metadata = document_dict.get("struct_data", {})
- doc_metadata["id"] = document_dict["id"]
-
- if chunk_type not in derived_struct_data:
- continue
-
- for chunk in derived_struct_data[chunk_type]:
- chunk_metadata = doc_metadata.copy()
- chunk_metadata["source"] = derived_struct_data.get("link", "")
-
- if chunk_type == "extractive_answers":
- chunk_metadata["source"] += f":{chunk.get('pageNumber', '')}"
-
- documents.append(
- Document(
- page_content=chunk.get("content", ""), metadata=chunk_metadata
- )
- )
-
- return documents
-
- def _convert_website_search_response(
- self, results: Sequence[SearchResult], chunk_type: str
- ) -> List[Document]:
- """Converts a sequence of search results to a list of LangChain documents."""
- from google.protobuf.json_format import MessageToDict
-
- documents: List[Document] = []
-
- for result in results:
- document_dict = MessageToDict(
- result.document._pb, preserving_proto_field_name=True
- )
- derived_struct_data = document_dict.get("derived_struct_data")
- if not derived_struct_data:
- continue
-
- doc_metadata = document_dict.get("struct_data", {})
- doc_metadata["id"] = document_dict["id"]
- doc_metadata["source"] = derived_struct_data.get("link", "")
- if derived_struct_data.get("title") is not None:
- doc_metadata["title"] = derived_struct_data.get("title")
-
- if chunk_type not in derived_struct_data:
- continue
-
- text_field = "snippet" if chunk_type == "snippets" else "content"
-
- for chunk in derived_struct_data[chunk_type]:
- documents.append(
- Document(
- page_content=chunk.get(text_field, ""), metadata=doc_metadata
- )
- )
-
- if not documents:
- print(f"No {chunk_type} could be found.") # noqa: T201
- if chunk_type == "extractive_answers":
- print( # noqa: T201
- "Make sure that your data store is using Advanced Website "
- "Indexing.\n"
- "https://cloud.google.com/generative-ai-app-builder/docs/about-advanced-features#advanced-website-indexing"
- )
-
- return documents
-
-
-@deprecated(
- since="0.0.33",
- removal="1.0",
- alternative_import="langchain_google_community.VertexAISearchRetriever",
-)
-class GoogleVertexAISearchRetriever(BaseRetriever, _BaseGoogleVertexAISearchRetriever):
- """`Google Vertex AI Search` retriever.
-
- For a detailed explanation of the Vertex AI Search concepts
- and configuration parameters, refer to the product documentation.
- https://cloud.google.com/generative-ai-app-builder/docs/enterprise-search-introduction
- """
-
- filter: Optional[str] = None
- """Filter expression."""
- get_extractive_answers: bool = False
- """If True return Extractive Answers, otherwise return Extractive Segments or Snippets.""" # noqa: E501
- max_documents: int = Field(default=5, ge=1, le=100)
- """The maximum number of documents to return."""
- max_extractive_answer_count: int = Field(default=1, ge=1, le=5)
- """The maximum number of extractive answers returned in each search result.
- At most 5 answers will be returned for each SearchResult.
- """
- max_extractive_segment_count: int = Field(default=1, ge=1, le=1)
- """The maximum number of extractive segments returned in each search result.
- Currently one segment will be returned for each SearchResult.
- """
- query_expansion_condition: int = Field(default=1, ge=0, le=2)
- """Specification to determine under which conditions query expansion should occur.
- 0 - Unspecified query expansion condition. In this case, server behavior defaults
- to disabled
- 1 - Disabled query expansion. Only the exact search query is used, even if
- SearchResponse.total_size is zero.
- 2 - Automatic query expansion built by the Search API.
- """
- spell_correction_mode: int = Field(default=2, ge=0, le=2)
- """Specification to determine under which conditions query expansion should occur.
- 0 - Unspecified spell correction mode. In this case, server behavior defaults
- to auto.
- 1 - Suggestion only. Search API will try to find a spell suggestion if there is any
- and put in the `SearchResponse.corrected_query`.
- The spell suggestion will not be used as the search query.
- 2 - Automatic spell correction built by the Search API.
- Search will be based on the corrected query if found.
- """
-
- # type is SearchServiceClient but can't be set due to optional imports
- _client: Any = None
- _serving_config: str
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- extra="ignore",
- )
-
- def __init__(self, **kwargs: Any) -> None:
- """Initializes private fields."""
- try:
- from google.cloud.discoveryengine_v1beta import SearchServiceClient
- except ImportError as exc:
- raise ImportError(
- "google.cloud.discoveryengine is not installed."
- "Please install it with pip install google-cloud-discoveryengine"
- ) from exc
-
- super().__init__(**kwargs)
-
- # For more information, refer to:
- # https://cloud.google.com/generative-ai-app-builder/docs/locations#specify_a_multi-region_for_your_data_store
- self._client = SearchServiceClient(
- credentials=self.credentials,
- client_options=self.client_options,
- client_info=get_client_info(module="vertex-ai-search"),
- )
-
- if self.engine_data_type == 3 and not self.search_engine_id:
- raise ValueError(
- "search_engine_id must be specified for blended search apps."
- )
-
- if self.search_engine_id:
- self._serving_config = f"projects/{self.project_id}/locations/{self.location_id}/collections/default_collection/engines/{self.search_engine_id}/servingConfigs/default_config" # noqa: E501
- elif self.data_store_id:
- self._serving_config = self._client.serving_config_path(
- project=self.project_id,
- location=self.location_id,
- data_store=self.data_store_id,
- serving_config=self.serving_config_id,
- )
- else:
- raise ValueError(
- "Either data_store_id or search_engine_id must be specified."
- )
-
- def _create_search_request(self, query: str) -> SearchRequest:
- """Prepares a SearchRequest object."""
- from google.cloud.discoveryengine_v1beta import SearchRequest
-
- query_expansion_spec = SearchRequest.QueryExpansionSpec(
- condition=self.query_expansion_condition,
- )
-
- spell_correction_spec = SearchRequest.SpellCorrectionSpec(
- mode=self.spell_correction_mode
- )
-
- if self.engine_data_type == 0:
- if self.get_extractive_answers:
- extractive_content_spec = (
- SearchRequest.ContentSearchSpec.ExtractiveContentSpec(
- max_extractive_answer_count=self.max_extractive_answer_count,
- )
- )
- else:
- extractive_content_spec = (
- SearchRequest.ContentSearchSpec.ExtractiveContentSpec(
- max_extractive_segment_count=self.max_extractive_segment_count,
- )
- )
- content_search_spec = SearchRequest.ContentSearchSpec(
- extractive_content_spec=extractive_content_spec
- )
- elif self.engine_data_type == 1:
- content_search_spec = None
- elif self.engine_data_type in (2, 3):
- content_search_spec = SearchRequest.ContentSearchSpec(
- extractive_content_spec=SearchRequest.ContentSearchSpec.ExtractiveContentSpec(
- max_extractive_answer_count=self.max_extractive_answer_count,
- ),
- snippet_spec=SearchRequest.ContentSearchSpec.SnippetSpec(
- return_snippet=True
- ),
- )
- else:
- raise NotImplementedError(
- "Only data store type 0 (Unstructured), 1 (Structured),"
- "2 (Website), or 3 (Blended) are supported currently."
- + f" Got {self.engine_data_type}"
- )
-
- return SearchRequest(
- query=query,
- filter=self.filter,
- serving_config=self._serving_config,
- page_size=self.max_documents,
- content_search_spec=content_search_spec,
- query_expansion_spec=query_expansion_spec,
- spell_correction_spec=spell_correction_spec,
- )
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- """Get documents relevant for a query."""
- return self.get_relevant_documents_with_response(query)[0]
-
- def get_relevant_documents_with_response(
- self, query: str
- ) -> Tuple[List[Document], Any]:
- from google.api_core.exceptions import InvalidArgument
-
- search_request = self._create_search_request(query)
-
- try:
- response = self._client.search(search_request)
- except InvalidArgument as exc:
- raise type(exc)(
- exc.message
- + " This might be due to engine_data_type not set correctly."
- )
-
- if self.engine_data_type == 0:
- chunk_type = (
- "extractive_answers"
- if self.get_extractive_answers
- else "extractive_segments"
- )
- documents = self._convert_unstructured_search_response(
- response.results, chunk_type
- )
- elif self.engine_data_type == 1:
- documents = self._convert_structured_search_response(response.results)
- elif self.engine_data_type in (2, 3):
- chunk_type = (
- "extractive_answers" if self.get_extractive_answers else "snippets"
- )
- documents = self._convert_website_search_response(
- response.results, chunk_type
- )
- else:
- raise NotImplementedError(
- "Only data store type 0 (Unstructured), 1 (Structured),"
- "2 (Website), or 3 (Blended) are supported currently."
- + f" Got {self.engine_data_type}"
- )
-
- return documents, response
-
-
-@deprecated(
- since="0.0.33",
- removal="1.0",
- alternative_import="langchain_google_community.VertexAIMultiTurnSearchRetriever",
-)
-class GoogleVertexAIMultiTurnSearchRetriever(
- BaseRetriever, _BaseGoogleVertexAISearchRetriever
-):
- """`Google Vertex AI Search` retriever for multi-turn conversations."""
-
- conversation_id: str = "-"
- """Vertex AI Search Conversation ID."""
-
- # type is ConversationalSearchServiceClient but can't be set due to optional imports
- _client: Any = None
- _serving_config: str
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- extra="ignore",
- )
-
- def __init__(self, **kwargs: Any):
- super().__init__(**kwargs)
- from google.cloud.discoveryengine_v1beta import (
- ConversationalSearchServiceClient,
- )
-
- self._client = ConversationalSearchServiceClient(
- credentials=self.credentials,
- client_options=self.client_options,
- client_info=get_client_info(module="vertex-ai-search"),
- )
-
- if not self.data_store_id:
- raise ValueError("data_store_id is required for MultiTurnSearchRetriever.")
-
- self._serving_config = self._client.serving_config_path(
- project=self.project_id,
- location=self.location_id,
- data_store=self.data_store_id,
- serving_config=self.serving_config_id,
- )
-
- if self.engine_data_type == 1 or self.engine_data_type == 3:
- raise NotImplementedError(
- "Data store type 1 (Structured) and 3 (Blended)"
- "is not currently supported for multi-turn search."
- + f" Got {self.engine_data_type}"
- )
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- """Get documents relevant for a query."""
- from google.cloud.discoveryengine_v1beta import (
- ConverseConversationRequest,
- TextInput,
- )
-
- request = ConverseConversationRequest(
- name=self._client.conversation_path(
- self.project_id,
- self.location_id,
- self.data_store_id,
- self.conversation_id,
- ),
- serving_config=self._serving_config,
- query=TextInput(input=query),
- )
- response = self._client.converse_conversation(request)
-
- if self.engine_data_type == 2:
- return self._convert_website_search_response(
- response.search_results, "extractive_answers"
- )
-
- return self._convert_unstructured_search_response(
- response.search_results, "extractive_answers"
- )
-
-
-class GoogleCloudEnterpriseSearchRetriever(GoogleVertexAISearchRetriever):
- """`Google Vertex Search API` retriever alias for backwards compatibility.
- DEPRECATED: Use `GoogleVertexAISearchRetriever` instead.
- """
-
- def __init__(self, **data: Any):
- import warnings
-
- warnings.warn(
- "GoogleCloudEnterpriseSearchRetriever is deprecated, use GoogleVertexAISearchRetriever", # noqa: E501
- DeprecationWarning,
- )
-
- super().__init__(**data)
diff --git a/libs/community/langchain_community/retrievers/kay.py b/libs/community/langchain_community/retrievers/kay.py
deleted file mode 100644
index ef594157b1..0000000000
--- a/libs/community/langchain_community/retrievers/kay.py
+++ /dev/null
@@ -1,60 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, List
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-
-class KayAiRetriever(BaseRetriever):
- """
- Retriever for Kay.ai datasets.
-
- To work properly, expects you to have KAY_API_KEY env variable set.
- You can get one for free at https://kay.ai/.
- """
-
- client: Any
- num_contexts: int
-
- @classmethod
- def create(
- cls,
- dataset_id: str,
- data_types: List[str],
- num_contexts: int = 6,
- ) -> KayAiRetriever:
- """
- Create a KayRetriever given a Kay dataset id and a list of datasources.
-
- Args:
- dataset_id: A dataset id category in Kay, like "company"
- data_types: A list of datasources present within a dataset. For
- "company" the corresponding datasources could be
- ["10-K", "10-Q", "8-K", "PressRelease"].
- num_contexts: The number of documents to retrieve on each query.
- Defaults to 6.
- """
- try:
- from kay.rag.retrievers import KayRetriever
- except ImportError:
- raise ImportError(
- "Could not import kay python package. Please install it with "
- "`pip install kay`.",
- )
-
- client = KayRetriever(dataset_id, data_types)
- return cls(client=client, num_contexts=num_contexts)
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- ctxs = self.client.query(query=query, num_context=self.num_contexts)
- docs = []
- for ctx in ctxs:
- page_content = ctx.pop("chunk_embed_text", None)
- if page_content is None:
- continue
- docs.append(Document(page_content=page_content, metadata={**ctx}))
- return docs
diff --git a/libs/community/langchain_community/retrievers/kendra.py b/libs/community/langchain_community/retrievers/kendra.py
deleted file mode 100644
index 899a0052e3..0000000000
--- a/libs/community/langchain_community/retrievers/kendra.py
+++ /dev/null
@@ -1,496 +0,0 @@
-import re
-from abc import ABC, abstractmethod
-from typing import (
- Any,
- Callable,
- Dict,
- List,
- Literal,
- Optional,
- Sequence,
- Union,
-)
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from pydantic import (
- BaseModel,
- Field,
- model_validator,
- validator,
-)
-from typing_extensions import Annotated
-
-
-def clean_excerpt(excerpt: str) -> str:
- """Clean an excerpt from Kendra.
-
- Args:
- excerpt: The excerpt to clean.
-
- Returns:
- The cleaned excerpt.
-
- """
- if not excerpt:
- return excerpt
- res = re.sub(r"\s+", " ", excerpt).replace("...", "")
- return res
-
-
-def combined_text(item: "ResultItem") -> str:
- """Combine a ResultItem title and excerpt into a single string.
-
- Args:
- item: the ResultItem of a Kendra search.
-
- Returns:
- A combined text of the title and excerpt of the given item.
-
- """
- text = ""
- title = item.get_title()
- if title:
- text += f"Document Title: {title}\n"
- excerpt = clean_excerpt(item.get_excerpt())
- if excerpt:
- text += f"Document Excerpt: \n{excerpt}\n"
- return text
-
-
-DocumentAttributeValueType = Union[str, int, List[str], None]
-"""Possible types of a DocumentAttributeValue.
-
-Dates are also represented as str.
-"""
-
-
-# Unexpected keyword argument "extra" for "__init_subclass__" of "object"
-class Highlight(BaseModel, extra="allow"):
- """Information that highlights the keywords in the excerpt."""
-
- BeginOffset: int
- """The zero-based location in the excerpt where the highlight starts."""
- EndOffset: int
- """The zero-based location in the excerpt where the highlight ends."""
- TopAnswer: Optional[bool]
- """Indicates whether the result is the best one."""
- Type: Optional[str]
- """The highlight type: STANDARD or THESAURUS_SYNONYM."""
-
-
-# Unexpected keyword argument "extra" for "__init_subclass__" of "object"
-class TextWithHighLights(BaseModel, extra="allow"):
- """Text with highlights."""
-
- Text: str
- """The text."""
- Highlights: Optional[Any]
- """The highlights."""
-
-
-# Unexpected keyword argument "extra" for "__init_subclass__" of "object"
-class AdditionalResultAttributeValue(BaseModel, extra="allow"):
- """Value of an additional result attribute."""
-
- TextWithHighlightsValue: TextWithHighLights
- """The text with highlights value."""
-
-
-# Unexpected keyword argument "extra" for "__init_subclass__" of "object"
-class AdditionalResultAttribute(BaseModel, extra="allow"):
- """Additional result attribute."""
-
- Key: str
- """The key of the attribute."""
- ValueType: Literal["TEXT_WITH_HIGHLIGHTS_VALUE"]
- """The type of the value."""
- Value: AdditionalResultAttributeValue
- """The value of the attribute."""
-
- def get_value_text(self) -> str:
- return self.Value.TextWithHighlightsValue.Text
-
-
-# Unexpected keyword argument "extra" for "__init_subclass__" of "object"
-class DocumentAttributeValue(BaseModel, extra="allow"):
- """Value of a document attribute."""
-
- DateValue: Optional[str] = None
- """The date expressed as an ISO 8601 string."""
- LongValue: Optional[int] = None
- """The long value."""
- StringListValue: Optional[List[str]] = None
- """The string list value."""
- StringValue: Optional[str] = None
- """The string value."""
-
- @property
- def value(self) -> DocumentAttributeValueType:
- """The only defined document attribute value or None.
- According to Amazon Kendra, you can only provide one
- value for a document attribute.
- """
- if self.DateValue:
- return self.DateValue
- if self.LongValue:
- return self.LongValue
- if self.StringListValue:
- return self.StringListValue
- if self.StringValue:
- return self.StringValue
-
- return None
-
-
-# Unexpected keyword argument "extra" for "__init_subclass__" of "object"
-class DocumentAttribute(BaseModel, extra="allow"):
- """Document attribute."""
-
- Key: str
- """The key of the attribute."""
- Value: DocumentAttributeValue
- """The value of the attribute."""
-
-
-# Unexpected keyword argument "extra" for "__init_subclass__" of "object"
-class ResultItem(BaseModel, ABC, extra="allow"):
- """Base class of a result item."""
-
- Id: Optional[str]
- """The ID of the relevant result item."""
- DocumentId: Optional[str]
- """The document ID."""
- DocumentURI: Optional[str]
- """The document URI."""
- DocumentAttributes: Optional[List[DocumentAttribute]] = []
- """The document attributes."""
- ScoreAttributes: Optional[dict]
- """The kendra score confidence"""
-
- @abstractmethod
- def get_title(self) -> str:
- """Document title."""
-
- @abstractmethod
- def get_excerpt(self) -> str:
- """Document excerpt or passage original content as retrieved by Kendra."""
-
- def get_additional_metadata(self) -> dict:
- """Document additional metadata dict.
- This returns any extra metadata except these:
- * result_id
- * document_id
- * source
- * title
- * excerpt
- * document_attributes
- """
- return {}
-
- def get_document_attributes_dict(self) -> Dict[str, DocumentAttributeValueType]:
- """Document attributes dict."""
- return {attr.Key: attr.Value.value for attr in (self.DocumentAttributes or [])}
-
- def get_score_attribute(self) -> str:
- """Document Score Confidence"""
- if self.ScoreAttributes is not None:
- return self.ScoreAttributes["ScoreConfidence"]
- else:
- return "NOT_AVAILABLE"
-
- def to_doc(
- self, page_content_formatter: Callable[["ResultItem"], str] = combined_text
- ) -> Document:
- """Converts this item to a Document."""
- page_content = page_content_formatter(self)
- metadata = self.get_additional_metadata()
- metadata.update(
- {
- "result_id": self.Id,
- "document_id": self.DocumentId,
- "source": self.DocumentURI,
- "title": self.get_title(),
- "excerpt": self.get_excerpt(),
- "document_attributes": self.get_document_attributes_dict(),
- "score": self.get_score_attribute(),
- }
- )
- return Document(page_content=page_content, metadata=metadata)
-
-
-class QueryResultItem(ResultItem):
- """Query API result item."""
-
- DocumentTitle: TextWithHighLights
- """The document title."""
- FeedbackToken: Optional[str]
- """Identifies a particular result from a particular query."""
- Format: Optional[str]
- """
- If the Type is ANSWER, then format is either:
- * TABLE: a table excerpt is returned in TableExcerpt;
- * TEXT: a text excerpt is returned in DocumentExcerpt.
- """
- Type: Optional[str]
- """Type of result: DOCUMENT or QUESTION_ANSWER or ANSWER"""
- AdditionalAttributes: Optional[List[AdditionalResultAttribute]] = []
- """One or more additional attributes associated with the result."""
- DocumentExcerpt: Optional[TextWithHighLights]
- """Excerpt of the document text."""
-
- def get_title(self) -> str:
- return self.DocumentTitle.Text
-
- def get_attribute_value(self) -> str:
- if not self.AdditionalAttributes:
- return ""
- if not self.AdditionalAttributes[0]:
- return ""
- else:
- return self.AdditionalAttributes[0].get_value_text()
-
- def get_excerpt(self) -> str:
- if (
- self.AdditionalAttributes
- and self.AdditionalAttributes[0].Key == "AnswerText"
- ):
- excerpt = self.get_attribute_value()
- elif self.DocumentExcerpt:
- excerpt = self.DocumentExcerpt.Text
- else:
- excerpt = ""
-
- return excerpt
-
- def get_additional_metadata(self) -> dict:
- additional_metadata = {"type": self.Type}
- return additional_metadata
-
-
-class RetrieveResultItem(ResultItem):
- """Retrieve API result item."""
-
- DocumentTitle: Optional[str]
- """The document title."""
- Content: Optional[str]
- """The content of the item."""
-
- def get_title(self) -> str:
- return self.DocumentTitle or ""
-
- def get_excerpt(self) -> str:
- return self.Content or ""
-
-
-# Unexpected keyword argument "extra" for "__init_subclass__" of "object"
-class QueryResult(BaseModel, extra="allow"):
- """`Amazon Kendra Query API` search result.
-
- It is composed of:
- * Relevant suggested answers: either a text excerpt or table excerpt.
- * Matching FAQs or questions-answer from your FAQ file.
- * Documents including an excerpt of each document with its title.
- """
-
- ResultItems: List[QueryResultItem]
- """The result items."""
-
-
-# Unexpected keyword argument "extra" for "__init_subclass__" of "object"
-class RetrieveResult(BaseModel, extra="allow"):
- """`Amazon Kendra Retrieve API` search result.
-
- It is composed of:
- * relevant passages or text excerpts given an input query.
- """
-
- QueryId: str
- """The ID of the query."""
- ResultItems: List[RetrieveResultItem]
- """The result items."""
-
-
-KENDRA_CONFIDENCE_MAPPING = {
- "NOT_AVAILABLE": 0.0,
- "LOW": 0.25,
- "MEDIUM": 0.50,
- "HIGH": 0.75,
- "VERY_HIGH": 1.0,
-}
-
-
-@deprecated(
- since="0.3.16",
- removal="1.0",
- alternative_import="langchain_aws.AmazonKendraRetriever",
-)
-class AmazonKendraRetriever(BaseRetriever):
- """`Amazon Kendra Index` retriever.
-
- Args:
- index_id: Kendra index id
-
- region_name: The aws region e.g., `us-west-2`.
- Fallsback to AWS_DEFAULT_REGION env variable
- or region specified in ~/.aws/config.
-
- credentials_profile_name: The name of the profile in the ~/.aws/credentials
- or ~/.aws/config files, which has either access keys or role information
- specified. If not specified, the default credential profile or, if on an
- EC2 instance, credentials from IMDS will be used.
-
- top_k: No of results to return
-
- attribute_filter: Additional filtering of results based on metadata
- See: https://docs.aws.amazon.com/kendra/latest/APIReference
-
- document_relevance_override_configurations: Overrides relevance tuning
- configurations of fields/attributes set at the index level
- See: https://docs.aws.amazon.com/kendra/latest/APIReference
-
- page_content_formatter: generates the Document page_content
- allowing access to all result item attributes. By default, it uses
- the item's title and excerpt.
-
- client: boto3 client for Kendra
-
- user_context: Provides information about the user context
- See: https://docs.aws.amazon.com/kendra/latest/APIReference
-
- Example:
- .. code-block:: python
-
- retriever = AmazonKendraRetriever(
- index_id="c0806df7-e76b-4bce-9b5c-d5582f6b1a03"
- )
-
- """
-
- index_id: str
- region_name: Optional[str] = None
- credentials_profile_name: Optional[str] = None
- top_k: int = 3
- attribute_filter: Optional[Dict] = None
- document_relevance_override_configurations: Optional[List[Dict]] = None
- page_content_formatter: Callable[[ResultItem], str] = combined_text
- client: Any
- user_context: Optional[Dict] = None
- min_score_confidence: Annotated[Optional[float], Field(ge=0.0, le=1.0)]
-
- @validator("top_k")
- def validate_top_k(cls, value: int) -> int:
- if value < 0:
- raise ValueError(f"top_k ({value}) cannot be negative.")
- return value
-
- @model_validator(mode="before")
- @classmethod
- def create_client(cls, values: Dict[str, Any]) -> Any:
- top_k = values.get("top_k")
- if top_k is not None and top_k < 0:
- raise ValueError(f"top_k ({top_k}) cannot be negative.")
-
- if values.get("client") is not None:
- return values
-
- try:
- import boto3
-
- if values.get("credentials_profile_name"):
- session = boto3.Session(profile_name=values["credentials_profile_name"])
- else:
- # use default credentials
- session = boto3.Session()
-
- client_params = {}
- if values.get("region_name"):
- client_params["region_name"] = values["region_name"]
-
- values["client"] = session.client("kendra", **client_params)
-
- return values
- except ImportError:
- raise ImportError(
- "Could not import boto3 python package. "
- "Please install it with `pip install boto3`."
- )
- except Exception as e:
- raise ValueError(
- "Could not load credentials to authenticate with AWS client. "
- "Please check that credentials in the specified "
- "profile name are valid."
- ) from e
-
- def _kendra_query(self, query: str) -> Sequence[ResultItem]:
- kendra_kwargs = {
- "IndexId": self.index_id,
- # truncate the query to ensure that
- # there is no validation exception from Kendra.
- "QueryText": query.strip()[0:999],
- "PageSize": self.top_k,
- }
- if self.attribute_filter is not None:
- kendra_kwargs["AttributeFilter"] = self.attribute_filter
- if self.document_relevance_override_configurations is not None:
- kendra_kwargs["DocumentRelevanceOverrideConfigurations"] = (
- self.document_relevance_override_configurations
- )
- if self.user_context is not None:
- kendra_kwargs["UserContext"] = self.user_context
-
- response = self.client.retrieve(**kendra_kwargs)
- r_result = RetrieveResult.parse_obj(response)
- if r_result.ResultItems:
- return r_result.ResultItems
-
- # Retrieve API returned 0 results, fall back to Query API
- response = self.client.query(**kendra_kwargs)
- q_result = QueryResult.parse_obj(response)
- return q_result.ResultItems
-
- def _get_top_k_docs(self, result_items: Sequence[ResultItem]) -> List[Document]:
- top_docs = [
- item.to_doc(self.page_content_formatter)
- for item in result_items[: self.top_k]
- ]
- return top_docs
-
- def _filter_by_score_confidence(self, docs: List[Document]) -> List[Document]:
- """
- Filter out the records that have a score confidence
- greater than the required threshold.
- """
- if not self.min_score_confidence:
- return docs
- filtered_docs = [
- item
- for item in docs
- if (
- item.metadata.get("score") is not None
- and isinstance(item.metadata["score"], str)
- and KENDRA_CONFIDENCE_MAPPING.get(item.metadata["score"], 0.0)
- >= self.min_score_confidence
- )
- ]
- return filtered_docs
-
- def _get_relevant_documents(
- self,
- query: str,
- *,
- run_manager: CallbackManagerForRetrieverRun,
- ) -> List[Document]:
- """Run search on Kendra index and get top k documents
-
- Example:
- .. code-block:: python
-
- docs = retriever.invoke('This is my query')
-
- """
- result_items = self._kendra_query(query)
- top_k_docs = self._get_top_k_docs(result_items)
- return self._filter_by_score_confidence(top_k_docs)
diff --git a/libs/community/langchain_community/retrievers/knn.py b/libs/community/langchain_community/retrievers/knn.py
deleted file mode 100644
index 8c08479248..0000000000
--- a/libs/community/langchain_community/retrievers/knn.py
+++ /dev/null
@@ -1,107 +0,0 @@
-"""KNN Retriever.
-Largely based on
-https://github.com/karpathy/randomfun/blob/master/knn_vs_svm.ipynb"""
-
-from __future__ import annotations
-
-import concurrent.futures
-from typing import Any, Iterable, List, Optional
-
-import numpy as np
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.embeddings import Embeddings
-from langchain_core.retrievers import BaseRetriever
-from pydantic import ConfigDict
-
-
-def create_index(contexts: List[str], embeddings: Embeddings) -> np.ndarray:
- """
- Create an index of embeddings for a list of contexts.
-
- Args:
- contexts: List of contexts to embed.
- embeddings: Embeddings model to use.
-
- Returns:
- Index of embeddings.
- """
- with concurrent.futures.ThreadPoolExecutor() as executor:
- return np.array(list(executor.map(embeddings.embed_query, contexts)))
-
-
-class KNNRetriever(BaseRetriever):
- """`KNN` retriever."""
-
- embeddings: Embeddings
- """Embeddings model to use."""
- index: Any = None
- """Index of embeddings."""
- texts: List[str]
- """List of texts to index."""
- metadatas: Optional[List[dict]] = None
- """List of metadatas corresponding with each text."""
- k: int = 4
- """Number of results to return."""
- relevancy_threshold: Optional[float] = None
- """Threshold for relevancy."""
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- )
-
- @classmethod
- def from_texts(
- cls,
- texts: List[str],
- embeddings: Embeddings,
- metadatas: Optional[List[dict]] = None,
- **kwargs: Any,
- ) -> KNNRetriever:
- index = create_index(texts, embeddings)
- return cls(
- embeddings=embeddings,
- index=index,
- texts=texts,
- metadatas=metadatas,
- **kwargs,
- )
-
- @classmethod
- def from_documents(
- cls,
- documents: Iterable[Document],
- embeddings: Embeddings,
- **kwargs: Any,
- ) -> KNNRetriever:
- texts, metadatas = zip(*((d.page_content, d.metadata) for d in documents))
- return cls.from_texts(
- texts=texts, embeddings=embeddings, metadatas=metadatas, **kwargs
- )
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- query_embeds = np.array(self.embeddings.embed_query(query))
- # calc L2 norm
- index_embeds = self.index / np.sqrt((self.index**2).sum(1, keepdims=True))
- query_embeds = query_embeds / np.sqrt((query_embeds**2).sum())
-
- similarities = index_embeds.dot(query_embeds)
- sorted_ix = np.argsort(-similarities)
-
- denominator = np.max(similarities) - np.min(similarities) + 1e-6
- normalized_similarities = (similarities - np.min(similarities)) / denominator
-
- top_k_results = [
- Document(
- page_content=self.texts[row],
- metadata=self.metadatas[row] if self.metadatas else {},
- )
- for row in sorted_ix[0 : self.k]
- if (
- self.relevancy_threshold is None
- or normalized_similarities[row] >= self.relevancy_threshold
- )
- ]
- return top_k_results
diff --git a/libs/community/langchain_community/retrievers/llama_index.py b/libs/community/langchain_community/retrievers/llama_index.py
deleted file mode 100644
index 1ab75b572c..0000000000
--- a/libs/community/langchain_community/retrievers/llama_index.py
+++ /dev/null
@@ -1,86 +0,0 @@
-from typing import Any, Dict, List, cast
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from pydantic import Field
-
-
-class LlamaIndexRetriever(BaseRetriever):
- """`LlamaIndex` retriever.
-
- It is used for the question-answering with sources over
- an LlamaIndex data structure."""
-
- index: Any = None
- """LlamaIndex index to query."""
- query_kwargs: Dict = Field(default_factory=dict)
- """Keyword arguments to pass to the query method."""
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- """Get documents relevant for a query."""
- try:
- from llama_index.core.base.response.schema import Response
- from llama_index.core.indices.base import BaseGPTIndex
- except ImportError:
- raise ImportError(
- "You need to install `pip install llama-index` to use this retriever."
- )
- index = cast(BaseGPTIndex, self.index)
-
- response = index.query(query, **self.query_kwargs)
- response = cast(Response, response)
- # parse source nodes
- docs = []
- for source_node in response.source_nodes:
- metadata = source_node.metadata or {}
- docs.append(
- Document(page_content=source_node.get_content(), metadata=metadata)
- )
- return docs
-
-
-class LlamaIndexGraphRetriever(BaseRetriever):
- """`LlamaIndex` graph data structure retriever.
-
- It is used for question-answering with sources over an LlamaIndex
- graph data structure."""
-
- graph: Any = None
- """LlamaIndex graph to query."""
- query_configs: List[Dict] = Field(default_factory=list)
- """List of query configs to pass to the query method."""
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- """Get documents relevant for a query."""
- try:
- from llama_index.core.base.response.schema import Response
- from llama_index.core.composability.base import (
- QUERY_CONFIG_TYPE,
- ComposableGraph,
- )
- except ImportError:
- raise ImportError(
- "You need to install `pip install llama-index` to use this retriever."
- )
- graph = cast(ComposableGraph, self.graph)
-
- # for now, inject response_mode="no_text" into query configs
- for query_config in self.query_configs:
- query_config["response_mode"] = "no_text"
- query_configs = cast(List[QUERY_CONFIG_TYPE], self.query_configs)
- response = graph.query(query, query_configs=query_configs)
- response = cast(Response, response)
-
- # parse source nodes
- docs = []
- for source_node in response.source_nodes:
- metadata = source_node.metadata or {}
- docs.append(
- Document(page_content=source_node.get_content(), metadata=metadata)
- )
- return docs
diff --git a/libs/community/langchain_community/retrievers/metal.py b/libs/community/langchain_community/retrievers/metal.py
deleted file mode 100644
index df2f57f235..0000000000
--- a/libs/community/langchain_community/retrievers/metal.py
+++ /dev/null
@@ -1,43 +0,0 @@
-from typing import Any, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from pydantic import model_validator
-
-
-class MetalRetriever(BaseRetriever):
- """`Metal API` retriever."""
-
- client: Any
- """The Metal client to use."""
- params: Optional[dict] = None
- """The parameters to pass to the Metal client."""
-
- @model_validator(mode="before")
- @classmethod
- def validate_client(cls, values: dict) -> Any:
- """Validate that the client is of the correct type."""
- from metal_sdk.metal import Metal
-
- if "client" in values:
- client = values["client"]
- if not isinstance(client, Metal):
- raise ValueError(
- "Got unexpected client, should be of type metal_sdk.metal.Metal. "
- f"Instead, got {type(client)}"
- )
-
- values["params"] = values.get("params", {})
-
- return values
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- results = self.client.search({"text": query}, **self.params)
- final_results = []
- for r in results["data"]:
- metadata = {k: v for k, v in r.items() if k != "text"}
- final_results.append(Document(page_content=r["text"], metadata=metadata))
- return final_results
diff --git a/libs/community/langchain_community/retrievers/milvus.py b/libs/community/langchain_community/retrievers/milvus.py
deleted file mode 100644
index 1739dd83ec..0000000000
--- a/libs/community/langchain_community/retrievers/milvus.py
+++ /dev/null
@@ -1,150 +0,0 @@
-"""Milvus Retriever"""
-
-import warnings
-from typing import Any, Dict, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.embeddings import Embeddings
-from langchain_core.retrievers import BaseRetriever
-from pydantic import model_validator
-
-from langchain_community.vectorstores.milvus import Milvus
-
-# TODO: Update to MilvusClient + Hybrid Search when available
-
-
-class MilvusRetriever(BaseRetriever):
- """Milvus API retriever.
-
- See detailed instructions here: https://python.langchain.com/docs/integrations/retrievers/milvus_hybrid_search/
-
- Setup:
- Install ``langchain-milvus`` and other dependencies:
-
- .. code-block:: bash
-
- pip install -U pymilvus[model] langchain-milvus
-
- Key init args:
- collection: Milvus Collection
-
- Instantiate:
- .. code-block:: python
-
- retriever = MilvusCollectionHybridSearchRetriever(collection=collection)
-
- Usage:
- .. code-block:: python
-
- query = "What are the story about ventures?"
-
- retriever.invoke(query)
-
- .. code-block:: none
-
- [Document(page_content="In 'The Lost Expedition' by Caspian Grey...", metadata={'doc_id': '449281835035545843'}),
- Document(page_content="In 'The Phantom Pilgrim' by Rowan Welles...", metadata={'doc_id': '449281835035545845'}),
- Document(page_content="In 'The Dreamwalker's Journey' by Lyra Snow..", metadata={'doc_id': '449281835035545846'})]
-
- Use within a chain:
- .. code-block:: python
-
- from langchain_core.output_parsers import StrOutputParser
- from langchain_core.prompts import ChatPromptTemplate
- from langchain_core.runnables import RunnablePassthrough
- from langchain_openai import ChatOpenAI
-
- prompt = ChatPromptTemplate.from_template(
- \"\"\"Answer the question based only on the context provided.
-
- Context: {context}
-
- Question: {question}\"\"\"
- )
-
- llm = ChatOpenAI(model="gpt-3.5-turbo-0125")
-
- def format_docs(docs):
- return "\\n\\n".join(doc.page_content for doc in docs)
-
- chain = (
- {"context": retriever | format_docs, "question": RunnablePassthrough()}
- | prompt
- | llm
- | StrOutputParser()
- )
-
- chain.invoke("What novels has Lila written and what are their contents?")
-
- .. code-block:: none
-
- "Lila Rose has written 'The Memory Thief,' which follows a charismatic thief..."
-
- """ # noqa: E501
-
- embedding_function: Embeddings
- collection_name: str = "LangChainCollection"
- collection_properties: Optional[Dict[str, Any]] = None
- connection_args: Optional[Dict[str, Any]] = None
- consistency_level: str = "Session"
- search_params: Optional[dict] = None
-
- store: Milvus
- retriever: BaseRetriever
-
- @model_validator(mode="before")
- @classmethod
- def create_retriever(cls, values: Dict) -> Any:
- """Create the Milvus store and retriever."""
- values["store"] = Milvus(
- values["embedding_function"],
- values["collection_name"],
- values["collection_properties"],
- values["connection_args"],
- values["consistency_level"],
- )
- values["retriever"] = values["store"].as_retriever(
- search_kwargs={"param": values["search_params"]}
- )
- return values
-
- def add_texts(
- self, texts: List[str], metadatas: Optional[List[dict]] = None
- ) -> None:
- """Add text to the Milvus store
-
- Args:
- texts (List[str]): The text
- metadatas (List[dict]): Metadata dicts, must line up with existing store
- """
- self.store.add_texts(texts, metadatas)
-
- def _get_relevant_documents(
- self,
- query: str,
- *,
- run_manager: CallbackManagerForRetrieverRun,
- **kwargs: Any,
- ) -> List[Document]:
- return self.retriever.invoke(
- query, run_manager=run_manager.get_child(), **kwargs
- )
-
-
-def MilvusRetreiver(*args: Any, **kwargs: Any) -> MilvusRetriever:
- """Deprecated MilvusRetreiver. Please use MilvusRetriever ('i' before 'e') instead.
-
- Args:
- *args:
- **kwargs:
-
- Returns:
- MilvusRetriever
- """
- warnings.warn(
- "MilvusRetreiver will be deprecated in the future. "
- "Please use MilvusRetriever ('i' before 'e') instead.",
- DeprecationWarning,
- )
- return MilvusRetriever(*args, **kwargs)
diff --git a/libs/community/langchain_community/retrievers/nanopq.py b/libs/community/langchain_community/retrievers/nanopq.py
deleted file mode 100644
index 274ad4b42e..0000000000
--- a/libs/community/langchain_community/retrievers/nanopq.py
+++ /dev/null
@@ -1,125 +0,0 @@
-from __future__ import annotations
-
-import concurrent.futures
-from typing import Any, Iterable, List, Optional
-
-import numpy as np
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.embeddings import Embeddings
-from langchain_core.retrievers import BaseRetriever
-from pydantic import ConfigDict
-
-
-def create_index(contexts: List[str], embeddings: Embeddings) -> np.ndarray:
- """
- Create an index of embeddings for a list of contexts.
-
- Args:
- contexts: List of contexts to embed.
- embeddings: Embeddings model to use.
-
- Returns:
- Index of embeddings.
- """
- with concurrent.futures.ThreadPoolExecutor() as executor:
- return np.array(list(executor.map(embeddings.embed_query, contexts)))
-
-
-class NanoPQRetriever(BaseRetriever):
- """`NanoPQ retriever."""
-
- embeddings: Embeddings
- """Embeddings model to use."""
- index: Any = None
- """Index of embeddings."""
- texts: List[str]
- """List of texts to index."""
- metadatas: Optional[List[dict]] = None
- """List of metadatas corresponding with each text."""
- k: int = 4
- """Number of results to return."""
- relevancy_threshold: Optional[float] = None
- """Threshold for relevancy."""
- subspace: int = 4
- """No of subspaces to be created, should be a multiple of embedding shape"""
- clusters: int = 128
- """No of clusters to be created"""
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- )
-
- @classmethod
- def from_texts(
- cls,
- texts: List[str],
- embeddings: Embeddings,
- metadatas: Optional[List[dict]] = None,
- **kwargs: Any,
- ) -> NanoPQRetriever:
- index = create_index(texts, embeddings)
- return cls(
- embeddings=embeddings,
- index=index,
- texts=texts,
- metadatas=metadatas,
- **kwargs,
- )
-
- @classmethod
- def from_documents(
- cls,
- documents: Iterable[Document],
- embeddings: Embeddings,
- **kwargs: Any,
- ) -> NanoPQRetriever:
- texts, metadatas = zip(*((d.page_content, d.metadata) for d in documents))
- return cls.from_texts(
- texts=texts, embeddings=embeddings, metadatas=metadatas, **kwargs
- )
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- try:
- from nanopq import PQ
- except ImportError:
- raise ImportError(
- "Could not import nanopq, please install with `pip install nanopq`."
- )
-
- query_embeds = np.array(self.embeddings.embed_query(query))
- try:
- pq = PQ(M=self.subspace, Ks=self.clusters, verbose=True).fit(
- self.index.astype("float32")
- )
- except AssertionError:
- error_message = (
- "Received params: training_sample={training_sample}, "
- "n_cluster={n_clusters}, subspace={subspace}, "
- "embedding_shape={embedding_shape}. Issue with the combination. "
- "Please retrace back to find the exact error"
- ).format(
- training_sample=self.index.shape[0],
- n_clusters=self.clusters,
- subspace=self.subspace,
- embedding_shape=self.index.shape[1],
- )
- raise RuntimeError(error_message)
-
- index_code = pq.encode(vecs=self.index.astype("float32"))
- dt = pq.dtable(query=query_embeds.astype("float32"))
- dists = dt.adist(codes=index_code)
-
- sorted_ix = np.argsort(dists)
-
- top_k_results = [
- Document(
- page_content=self.texts[row],
- metadata=self.metadatas[row] if self.metadatas else {},
- )
- for row in sorted_ix[0 : self.k]
- ]
-
- return top_k_results
diff --git a/libs/community/langchain_community/retrievers/needle.py b/libs/community/langchain_community/retrievers/needle.py
deleted file mode 100644
index 52a5245108..0000000000
--- a/libs/community/langchain_community/retrievers/needle.py
+++ /dev/null
@@ -1,101 +0,0 @@
-from typing import Any, List, Optional # noqa: I001
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from pydantic import BaseModel, Field
-
-
-class NeedleRetriever(BaseRetriever, BaseModel):
- """
- NeedleRetriever retrieves relevant documents or context from a Needle collection
- based on a search query.
-
- Setup:
- Install the `needle-python` library and set your Needle API key.
-
- .. code-block:: bash
-
- pip install needle-python
- export NEEDLE_API_KEY="your-api-key"
-
- Key init args:
- - `needle_api_key` (Optional[str]): The API key for authenticating with Needle.
- - `collection_id` (str): The ID of the Needle collection to search in.
- - `client` (Optional[NeedleClient]): An optional instance of the NeedleClient.
- - `top_k` (Optional[int]): Maximum number of results to return.
-
- Usage:
- .. code-block:: python
-
- from langchain_community.retrievers.needle import NeedleRetriever
-
- retriever = NeedleRetriever(
- needle_api_key="your-api-key",
- collection_id="your-collection-id",
- top_k=10 # optional
- )
-
- results = retriever.retrieve("example query")
- for doc in results:
- print(doc.page_content)
- """
-
- client: Optional[Any] = None
- """Optional instance of NeedleClient."""
- needle_api_key: Optional[str] = Field(None, description="Needle API Key")
- collection_id: Optional[str] = Field(
- ..., description="The ID of the Needle collection to search in"
- )
- top_k: Optional[int] = Field(
- default=None, description="Maximum number of search results to return"
- )
-
- def _initialize_client(self) -> None:
- """
- Initialize the NeedleClient with the provided API key.
-
- If a client instance is already provided, this method does nothing.
- """
- try:
- from needle.v1 import NeedleClient
- except ImportError:
- raise ImportError("Please install with `pip install needle-python`.")
-
- if not self.client:
- self.client = NeedleClient(api_key=self.needle_api_key)
-
- def _search_collection(self, query: str) -> List[Document]:
- """
- Search the Needle collection for relevant documents.
-
- Args:
- query (str): The search query used to find relevant documents.
-
- Returns:
- List[Document]: A list of documents matching the search query.
- """
- self._initialize_client()
- if self.client is None:
- raise ValueError("NeedleClient is not initialized. Provide an API key.")
-
- results = self.client.collections.search(
- collection_id=self.collection_id, text=query, top_k=self.top_k
- )
- docs = [Document(page_content=result.content) for result in results]
- return docs
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- """
- Retrieve relevant documents based on the query.
-
- Args:
- query (str): The query string used to search the collection.
- Returns:
- List[Document]: A list of documents relevant to the query.
- """
- # The `run_manager` parameter is included to match the superclass signature,
- # but it is not used in this implementation.
- return self._search_collection(query)
diff --git a/libs/community/langchain_community/retrievers/outline.py b/libs/community/langchain_community/retrievers/outline.py
deleted file mode 100644
index 03b1118125..0000000000
--- a/libs/community/langchain_community/retrievers/outline.py
+++ /dev/null
@@ -1,20 +0,0 @@
-from typing import List
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-from langchain_community.utilities.outline import OutlineAPIWrapper
-
-
-class OutlineRetriever(BaseRetriever, OutlineAPIWrapper):
- """Retriever for Outline API.
-
- It wraps run() to get_relevant_documents().
- It uses all OutlineAPIWrapper arguments without any change.
- """
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- return self.run(query=query)
diff --git a/libs/community/langchain_community/retrievers/pinecone_hybrid_search.py b/libs/community/langchain_community/retrievers/pinecone_hybrid_search.py
deleted file mode 100644
index cd3e3e96d0..0000000000
--- a/libs/community/langchain_community/retrievers/pinecone_hybrid_search.py
+++ /dev/null
@@ -1,185 +0,0 @@
-"""Taken from: https://docs.pinecone.io/docs/hybrid-search"""
-
-import hashlib
-from typing import Any, Dict, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.embeddings import Embeddings
-from langchain_core.retrievers import BaseRetriever
-from langchain_core.utils import pre_init
-from pydantic import ConfigDict
-
-
-def hash_text(text: str) -> str:
- """Hash a text using SHA256.
-
- Args:
- text: Text to hash.
-
- Returns:
- Hashed text.
- """
- return str(hashlib.sha256(text.encode("utf-8")).hexdigest())
-
-
-def create_index(
- contexts: List[str],
- index: Any,
- embeddings: Embeddings,
- sparse_encoder: Any,
- ids: Optional[List[str]] = None,
- metadatas: Optional[List[dict]] = None,
- namespace: Optional[str] = None,
- text_key: str = "context",
-) -> None:
- """Create an index from a list of contexts.
-
- It modifies the index argument in-place!
-
- Args:
- contexts: List of contexts to embed.
- index: Index to use.
- embeddings: Embeddings model to use.
- sparse_encoder: Sparse encoder to use.
- ids: List of ids to use for the documents.
- metadatas: List of metadata to use for the documents.
- namespace: Namespace value for index partition.
- """
- batch_size = 32
- _iterator = range(0, len(contexts), batch_size)
- try:
- from tqdm.auto import tqdm
-
- _iterator = tqdm(_iterator)
- except ImportError:
- pass
-
- if ids is None:
- # create unique ids using hash of the text
- ids = [hash_text(context) for context in contexts]
-
- for i in _iterator:
- # find end of batch
- i_end = min(i + batch_size, len(contexts))
- # extract batch
- context_batch = contexts[i:i_end]
- batch_ids = ids[i:i_end]
- metadata_batch = (
- metadatas[i:i_end] if metadatas else [{} for _ in context_batch]
- )
- # add context passages as metadata
- meta = [
- {text_key: context, **metadata}
- for context, metadata in zip(context_batch, metadata_batch)
- ]
-
- # create dense vectors
- dense_embeds = embeddings.embed_documents(context_batch)
- # create sparse vectors
- sparse_embeds = sparse_encoder.encode_documents(context_batch)
- for s in sparse_embeds:
- s["values"] = [float(s1) for s1 in s["values"]]
-
- vectors = []
- # loop through the data and create dictionaries for upserts
- for doc_id, sparse, dense, metadata in zip(
- batch_ids, sparse_embeds, dense_embeds, meta
- ):
- vectors.append(
- {
- "id": doc_id,
- "sparse_values": sparse,
- "values": dense,
- "metadata": metadata,
- }
- )
-
- # upload the documents to the new hybrid index
- index.upsert(vectors, namespace=namespace)
-
-
-class PineconeHybridSearchRetriever(BaseRetriever):
- """`Pinecone Hybrid Search` retriever."""
-
- embeddings: Embeddings
- """Embeddings model to use."""
- """description"""
- sparse_encoder: Any = None
- """Sparse encoder to use."""
- index: Any = None
- """Pinecone index to use."""
- top_k: int = 4
- """Number of documents to return."""
- alpha: float = 0.5
- """Alpha value for hybrid search."""
- namespace: Optional[str] = None
- """Namespace value for index partition."""
- text_key: str = "context"
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- extra="forbid",
- )
-
- def add_texts(
- self,
- texts: List[str],
- ids: Optional[List[str]] = None,
- metadatas: Optional[List[dict]] = None,
- namespace: Optional[str] = None,
- ) -> None:
- create_index(
- texts,
- self.index,
- self.embeddings,
- self.sparse_encoder,
- ids=ids,
- metadatas=metadatas,
- namespace=namespace,
- text_key=self.text_key,
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that api key and python package exists in environment."""
- try:
- from pinecone_text.hybrid import hybrid_convex_scale # noqa:F401
- from pinecone_text.sparse.base_sparse_encoder import (
- BaseSparseEncoder, # noqa:F401
- )
- except ImportError:
- raise ImportError(
- "Could not import pinecone_text python package. "
- "Please install it with `pip install pinecone_text`."
- )
- return values
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun, **kwargs: Any
- ) -> List[Document]:
- from pinecone_text.hybrid import hybrid_convex_scale
-
- sparse_vec = self.sparse_encoder.encode_queries(query)
- # convert the question into a dense vector
- dense_vec = self.embeddings.embed_query(query)
- # scale alpha with hybrid_scale
- dense_vec, sparse_vec = hybrid_convex_scale(dense_vec, sparse_vec, self.alpha)
- sparse_vec["values"] = [float(s1) for s1 in sparse_vec["values"]]
- # query pinecone with the query parameters
- result = self.index.query(
- vector=dense_vec,
- sparse_vector=sparse_vec,
- top_k=self.top_k,
- include_metadata=True,
- namespace=self.namespace,
- **kwargs,
- )
- final_result = []
- for res in result["matches"]:
- context = res["metadata"].pop(self.text_key)
- metadata = res["metadata"]
- if "score" not in metadata and "score" in res:
- metadata["score"] = res["score"]
- final_result.append(Document(page_content=context, metadata=metadata))
- # return search results as json
- return final_result
diff --git a/libs/community/langchain_community/retrievers/pubmed.py b/libs/community/langchain_community/retrievers/pubmed.py
deleted file mode 100644
index d68e85b80b..0000000000
--- a/libs/community/langchain_community/retrievers/pubmed.py
+++ /dev/null
@@ -1,20 +0,0 @@
-from typing import List
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-from langchain_community.utilities.pubmed import PubMedAPIWrapper
-
-
-class PubMedRetriever(BaseRetriever, PubMedAPIWrapper):
- """`PubMed API` retriever.
-
- It wraps load() to get_relevant_documents().
- It uses all PubMedAPIWrapper arguments without any change.
- """
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- return self.load_docs(query=query)
diff --git a/libs/community/langchain_community/retrievers/pupmed.py b/libs/community/langchain_community/retrievers/pupmed.py
deleted file mode 100644
index b4318034b2..0000000000
--- a/libs/community/langchain_community/retrievers/pupmed.py
+++ /dev/null
@@ -1,5 +0,0 @@
-from langchain_community.retrievers.pubmed import PubMedRetriever
-
-__all__ = [
- "PubMedRetriever",
-]
diff --git a/libs/community/langchain_community/retrievers/qdrant_sparse_vector_retriever.py b/libs/community/langchain_community/retrievers/qdrant_sparse_vector_retriever.py
deleted file mode 100644
index 1b64c3467f..0000000000
--- a/libs/community/langchain_community/retrievers/qdrant_sparse_vector_retriever.py
+++ /dev/null
@@ -1,220 +0,0 @@
-import uuid
-from itertools import islice
-from typing import (
- Any,
- Callable,
- Dict,
- Generator,
- Iterable,
- List,
- Optional,
- Sequence,
- Tuple,
- cast,
-)
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from langchain_core.utils import pre_init
-from pydantic import ConfigDict
-
-from langchain_community.vectorstores.qdrant import Qdrant, QdrantException
-
-
-@deprecated(
- since="0.2.16",
- alternative=(
- "Qdrant vector store now supports sparse retrievals natively. "
- "Use langchain_qdrant.QdrantVectorStore#as_retriever() instead. "
- "Reference: "
- "https://python.langchain.com/docs/integrations/vectorstores/qdrant/#sparse-vector-search"
- ),
- removal="0.5.0",
-)
-class QdrantSparseVectorRetriever(BaseRetriever):
- """Qdrant sparse vector retriever."""
-
- client: Any = None
- """'qdrant_client' instance to use."""
- collection_name: str
- """Qdrant collection name."""
- sparse_vector_name: str
- """Name of the sparse vector to use."""
- sparse_encoder: Callable[[str], Tuple[List[int], List[float]]]
- """Sparse encoder function to use."""
- k: int = 4
- """Number of documents to return per query. Defaults to 4."""
- filter: Optional[Any] = None
- """Qdrant qdrant_client.models.Filter to use for queries. Defaults to None."""
- content_payload_key: str = "content"
- """Payload field containing the document content. Defaults to 'content'"""
- metadata_payload_key: str = "metadata"
- """Payload field containing the document metadata. Defaults to 'metadata'."""
- search_options: Dict[str, Any] = {}
- """Additional search options to pass to qdrant_client.QdrantClient.search()."""
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- extra="forbid",
- )
-
- @pre_init
- def validate_environment(cls, values: Dict) -> Dict:
- """Validate that 'qdrant_client' python package exists in environment."""
- try:
- from grpc import RpcError
- from qdrant_client import QdrantClient, models
- from qdrant_client.http.exceptions import UnexpectedResponse
- except ImportError:
- raise ImportError(
- "Could not import qdrant-client python package. "
- "Please install it with `pip install qdrant-client`."
- )
-
- client = values["client"]
- if not isinstance(client, QdrantClient):
- raise ValueError(
- f"client should be an instance of qdrant_client.QdrantClient, "
- f"got {type(client)}"
- )
-
- filter = values["filter"]
- if filter is not None and not isinstance(filter, models.Filter):
- raise ValueError(
- f"filter should be an instance of qdrant_client.models.Filter, "
- f"got {type(filter)}"
- )
-
- client = cast(QdrantClient, client)
-
- collection_name = values["collection_name"]
- sparse_vector_name = values["sparse_vector_name"]
-
- try:
- collection_info = client.get_collection(collection_name)
- sparse_vectors_config = collection_info.config.params.sparse_vectors
-
- if sparse_vector_name not in sparse_vectors_config:
- raise QdrantException(
- f"Existing Qdrant collection {collection_name} does not "
- f"contain sparse vector named {sparse_vector_name}."
- f"Did you mean one of {', '.join(sparse_vectors_config.keys())}?"
- )
- except (UnexpectedResponse, RpcError, ValueError):
- raise QdrantException(
- f"Qdrant collection {collection_name} does not exist."
- )
- return values
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- from qdrant_client import QdrantClient, models
-
- client = cast(QdrantClient, self.client)
- query_indices, query_values = self.sparse_encoder(query)
- results = client.search(
- self.collection_name,
- query_filter=self.filter,
- query_vector=models.NamedSparseVector(
- name=self.sparse_vector_name,
- vector=models.SparseVector(
- indices=query_indices,
- values=query_values,
- ),
- ),
- limit=self.k,
- with_vectors=False,
- **self.search_options,
- )
- return [
- Qdrant._document_from_scored_point(
- point,
- self.collection_name,
- self.content_payload_key,
- self.metadata_payload_key,
- )
- for point in results
- ]
-
- def add_documents(self, documents: List[Document], **kwargs: Any) -> List[str]:
- """Run more documents through the embeddings and add to the vectorstore.
-
- Args:
- documents (List[Document]: Documents to add to the vectorstore.
-
- Returns:
- List[str]: List of IDs of the added texts.
- """
- texts = [doc.page_content for doc in documents]
- metadatas = [doc.metadata for doc in documents]
- return self.add_texts(texts, metadatas, **kwargs)
-
- def add_texts(
- self,
- texts: Iterable[str],
- metadatas: Optional[List[dict]] = None,
- ids: Optional[Sequence[str]] = None,
- batch_size: int = 64,
- **kwargs: Any,
- ) -> List[str]:
- from qdrant_client import QdrantClient
-
- added_ids = []
- client = cast(QdrantClient, self.client)
- for batch_ids, points in self._generate_rest_batches(
- texts, metadatas, ids, batch_size
- ):
- client.upsert(self.collection_name, points=points, **kwargs)
- added_ids.extend(batch_ids)
-
- return added_ids
-
- def _generate_rest_batches(
- self,
- texts: Iterable[str],
- metadatas: Optional[List[dict]] = None,
- ids: Optional[Sequence[str]] = None,
- batch_size: int = 64,
- ) -> Generator[Tuple[List[str], List[Any]], None, None]:
- from qdrant_client import models as rest
-
- texts_iterator = iter(texts)
- metadatas_iterator = iter(metadatas or [])
- ids_iterator = iter(ids or [uuid.uuid4().hex for _ in iter(texts)])
- while batch_texts := list(islice(texts_iterator, batch_size)):
- # Take the corresponding metadata and id for each text in a batch
- batch_metadatas = list(islice(metadatas_iterator, batch_size)) or None
- batch_ids = list(islice(ids_iterator, batch_size))
-
- # Generate the sparse embeddings for all the texts in a batch
- batch_embeddings: List[Tuple[List[int], List[float]]] = [
- self.sparse_encoder(text) for text in batch_texts
- ]
-
- points = [
- rest.PointStruct(
- id=point_id,
- vector={
- self.sparse_vector_name: rest.SparseVector(
- indices=sparse_vector[0],
- values=sparse_vector[1],
- )
- },
- payload=payload,
- )
- for point_id, sparse_vector, payload in zip(
- batch_ids,
- batch_embeddings,
- Qdrant._build_payloads(
- batch_texts,
- batch_metadatas,
- self.content_payload_key,
- self.metadata_payload_key,
- ),
- )
- ]
-
- yield batch_ids, points
diff --git a/libs/community/langchain_community/retrievers/rememberizer.py b/libs/community/langchain_community/retrievers/rememberizer.py
deleted file mode 100644
index c0aae8bd52..0000000000
--- a/libs/community/langchain_community/retrievers/rememberizer.py
+++ /dev/null
@@ -1,20 +0,0 @@
-from typing import List
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-from langchain_community.utilities.rememberizer import RememberizerAPIWrapper
-
-
-class RememberizerRetriever(BaseRetriever, RememberizerAPIWrapper):
- """`Rememberizer` retriever.
-
- It wraps load() to get_relevant_documents().
- It uses all RememberizerAPIWrapper arguments without any change.
- """
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- return self.load(query=query)
diff --git a/libs/community/langchain_community/retrievers/remote_retriever.py b/libs/community/langchain_community/retrievers/remote_retriever.py
deleted file mode 100644
index f384385557..0000000000
--- a/libs/community/langchain_community/retrievers/remote_retriever.py
+++ /dev/null
@@ -1,56 +0,0 @@
-from typing import List, Optional
-
-import aiohttp
-import requests
-from langchain_core.callbacks import (
- AsyncCallbackManagerForRetrieverRun,
- CallbackManagerForRetrieverRun,
-)
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-
-class RemoteLangChainRetriever(BaseRetriever):
- """`LangChain API` retriever."""
-
- url: str
- """URL of the remote LangChain API."""
- headers: Optional[dict] = None
- """Headers to use for the request."""
- input_key: str = "message"
- """Key to use for the input in the request."""
- response_key: str = "response"
- """Key to use for the response in the request."""
- page_content_key: str = "page_content"
- """Key to use for the page content in the response."""
- metadata_key: str = "metadata"
- """Key to use for the metadata in the response."""
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- response = requests.post(
- self.url, json={self.input_key: query}, headers=self.headers
- )
- result = response.json()
- return [
- Document(
- page_content=r[self.page_content_key], metadata=r[self.metadata_key]
- )
- for r in result[self.response_key]
- ]
-
- async def _aget_relevant_documents(
- self, query: str, *, run_manager: AsyncCallbackManagerForRetrieverRun
- ) -> List[Document]:
- async with aiohttp.ClientSession() as session:
- async with session.request(
- "POST", self.url, headers=self.headers, json={self.input_key: query}
- ) as response:
- result = await response.json()
- return [
- Document(
- page_content=r[self.page_content_key], metadata=r[self.metadata_key]
- )
- for r in result[self.response_key]
- ]
diff --git a/libs/community/langchain_community/retrievers/svm.py b/libs/community/langchain_community/retrievers/svm.py
deleted file mode 100644
index 58a7889691..0000000000
--- a/libs/community/langchain_community/retrievers/svm.py
+++ /dev/null
@@ -1,127 +0,0 @@
-from __future__ import annotations
-
-import concurrent.futures
-from typing import Any, Iterable, List, Optional
-
-import numpy as np
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.embeddings import Embeddings
-from langchain_core.retrievers import BaseRetriever
-from pydantic import ConfigDict
-
-
-def create_index(contexts: List[str], embeddings: Embeddings) -> np.ndarray:
- """
- Create an index of embeddings for a list of contexts.
-
- Args:
- contexts: List of contexts to embed.
- embeddings: Embeddings model to use.
-
- Returns:
- Index of embeddings.
- """
- with concurrent.futures.ThreadPoolExecutor() as executor:
- return np.array(list(executor.map(embeddings.embed_query, contexts)))
-
-
-class SVMRetriever(BaseRetriever):
- """`SVM` retriever.
-
- Largely based on
- https://github.com/karpathy/randomfun/blob/master/knn_vs_svm.ipynb
- """
-
- embeddings: Embeddings
- """Embeddings model to use."""
- index: Any = None
- """Index of embeddings."""
- texts: List[str]
- """List of texts to index."""
- metadatas: Optional[List[dict]] = None
- """List of metadatas corresponding with each text."""
- k: int = 4
- """Number of results to return."""
- relevancy_threshold: Optional[float] = None
- """Threshold for relevancy."""
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- )
-
- @classmethod
- def from_texts(
- cls,
- texts: List[str],
- embeddings: Embeddings,
- metadatas: Optional[List[dict]] = None,
- **kwargs: Any,
- ) -> SVMRetriever:
- index = create_index(texts, embeddings)
- return cls(
- embeddings=embeddings,
- index=index,
- texts=texts,
- metadatas=metadatas,
- **kwargs,
- )
-
- @classmethod
- def from_documents(
- cls,
- documents: Iterable[Document],
- embeddings: Embeddings,
- **kwargs: Any,
- ) -> SVMRetriever:
- texts, metadatas = zip(*((d.page_content, d.metadata) for d in documents))
- return cls.from_texts(
- texts=texts, embeddings=embeddings, metadatas=metadatas, **kwargs
- )
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- try:
- from sklearn import svm
- except ImportError:
- raise ImportError(
- "Could not import scikit-learn, please install with `pip install "
- "scikit-learn`."
- )
-
- query_embeds = np.array(self.embeddings.embed_query(query))
- x = np.concatenate([query_embeds[None, ...], self.index])
- y = np.zeros(x.shape[0])
- y[0] = 1
-
- clf = svm.LinearSVC(
- class_weight="balanced", verbose=False, max_iter=10000, tol=1e-6, C=0.1
- )
- clf.fit(x, y)
-
- similarities = clf.decision_function(x)
- sorted_ix = np.argsort(-similarities)
-
- # svm.LinearSVC in scikit-learn is non-deterministic.
- # if a text is the same as a query, there is no guarantee
- # the query will be in the first index.
- # this performs a simple swap, this works because anything
- # left of the 0 should be equivalent.
- zero_index = np.where(sorted_ix == 0)[0][0]
- if zero_index != 0:
- sorted_ix[0], sorted_ix[zero_index] = sorted_ix[zero_index], sorted_ix[0]
-
- denominator = np.max(similarities) - np.min(similarities) + 1e-6
- normalized_similarities = (similarities - np.min(similarities)) / denominator
-
- top_k_results = []
- for row in sorted_ix[1 : self.k + 1]:
- if (
- self.relevancy_threshold is None
- or normalized_similarities[row] >= self.relevancy_threshold
- ):
- metadata = self.metadatas[row - 1] if self.metadatas else {}
- doc = Document(page_content=self.texts[row - 1], metadata=metadata)
- top_k_results.append(doc)
- return top_k_results
diff --git a/libs/community/langchain_community/retrievers/tavily_search_api.py b/libs/community/langchain_community/retrievers/tavily_search_api.py
deleted file mode 100644
index aa0a08e3bf..0000000000
--- a/libs/community/langchain_community/retrievers/tavily_search_api.py
+++ /dev/null
@@ -1,152 +0,0 @@
-import os
-from enum import Enum
-from typing import Any, Dict, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-
-class SearchDepth(Enum):
- """Search depth as enumerator."""
-
- BASIC = "basic"
- ADVANCED = "advanced"
-
-
-class TavilySearchAPIRetriever(BaseRetriever):
- """Tavily Search API retriever.
-
- Setup:
- Install ``langchain-community`` and set environment variable ``TAVILY_API_KEY``.
-
- .. code-block:: bash
-
- pip install -U langchain-community
- export TAVILY_API_KEY="your-api-key"
-
- Key init args:
- k: int
- Number of results to include.
- include_generated_answer: bool
- Include a generated answer with results
- include_raw_content: bool
- Include raw content with results.
- include_images: bool
- Return images in addition to text.
-
- Instantiate:
- .. code-block:: python
-
- from langchain_community.retrievers import TavilySearchAPIRetriever
-
- retriever = TavilySearchAPIRetriever(k=3)
-
- Usage:
- .. code-block:: python
-
- query = "what year was breath of the wild released?"
-
- retriever.invoke(query)
-
- Use within a chain:
- .. code-block:: python
-
- from langchain_core.output_parsers import StrOutputParser
- from langchain_core.prompts import ChatPromptTemplate
- from langchain_core.runnables import RunnablePassthrough
- from langchain_openai import ChatOpenAI
-
- prompt = ChatPromptTemplate.from_template(
- \"\"\"Answer the question based only on the context provided.
-
- Context: {context}
-
- Question: {question}\"\"\"
- )
-
- llm = ChatOpenAI(model="gpt-3.5-turbo-0125")
-
- def format_docs(docs):
- return "\n\n".join(doc.page_content for doc in docs)
-
- chain = (
- {"context": retriever | format_docs, "question": RunnablePassthrough()}
- | prompt
- | llm
- | StrOutputParser()
- )
-
- chain.invoke("how many units did bretch of the wild sell in 2020")
-
- """ # noqa: E501
-
- k: int = 10
- include_generated_answer: bool = False
- include_raw_content: bool = False
- include_images: bool = False
- search_depth: SearchDepth = SearchDepth.BASIC
- include_domains: Optional[List[str]] = None
- exclude_domains: Optional[List[str]] = None
- kwargs: Optional[Dict[str, Any]] = {}
- api_key: Optional[str] = None
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- try:
- try:
- from tavily import TavilyClient
- except ImportError:
- # Older of tavily used Client
- from tavily import Client as TavilyClient
- except ImportError:
- raise ImportError(
- "Tavily python package not found. "
- "Please install it with `pip install tavily-python`."
- )
-
- tavily = TavilyClient(api_key=self.api_key or os.environ["TAVILY_API_KEY"])
- max_results = self.k if not self.include_generated_answer else self.k - 1
- response = tavily.search(
- query=query,
- max_results=max_results,
- search_depth=self.search_depth.value,
- include_answer=self.include_generated_answer,
- include_domains=self.include_domains,
- exclude_domains=self.exclude_domains,
- include_raw_content=self.include_raw_content,
- include_images=self.include_images,
- **self.kwargs,
- )
- docs = [
- Document(
- page_content=result.get("content", "")
- if not self.include_raw_content
- else (result.get("raw_content") or ""),
- metadata={
- "title": result.get("title", ""),
- "source": result.get("url", ""),
- **{
- k: v
- for k, v in result.items()
- if k not in ("content", "title", "url", "raw_content")
- },
- "images": response.get("images"),
- },
- )
- for result in response.get("results")
- ]
- if self.include_generated_answer:
- docs = [
- Document(
- page_content=response.get("answer", ""),
- metadata={
- "title": "Suggested Answer",
- "source": "https://tavily.com/",
- },
- ),
- *docs,
- ]
-
- return docs
diff --git a/libs/community/langchain_community/retrievers/tfidf.py b/libs/community/langchain_community/retrievers/tfidf.py
deleted file mode 100644
index 6a991f81f3..0000000000
--- a/libs/community/langchain_community/retrievers/tfidf.py
+++ /dev/null
@@ -1,159 +0,0 @@
-from __future__ import annotations
-
-import pickle
-from pathlib import Path
-from typing import Any, Dict, Iterable, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from pydantic import ConfigDict
-
-
-class TFIDFRetriever(BaseRetriever):
- """`TF-IDF` retriever.
-
- Largely based on
- https://github.com/asvskartheek/Text-Retrieval/blob/master/TF-IDF%20Search%20Engine%20(SKLEARN).ipynb
- """
-
- vectorizer: Any = None
- """TF-IDF vectorizer."""
- docs: List[Document]
- """Documents."""
- tfidf_array: Any = None
- """TF-IDF array."""
- k: int = 4
- """Number of documents to return."""
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- )
-
- @classmethod
- def from_texts(
- cls,
- texts: Iterable[str],
- metadatas: Optional[Iterable[dict]] = None,
- tfidf_params: Optional[Dict[str, Any]] = None,
- **kwargs: Any,
- ) -> TFIDFRetriever:
- try:
- from sklearn.feature_extraction.text import TfidfVectorizer
- except ImportError:
- raise ImportError(
- "Could not import scikit-learn, please install with `pip install "
- "scikit-learn`."
- )
-
- tfidf_params = tfidf_params or {}
- vectorizer = TfidfVectorizer(**tfidf_params)
- tfidf_array = vectorizer.fit_transform(texts)
- metadatas = metadatas or ({} for _ in texts)
- docs = [Document(page_content=t, metadata=m) for t, m in zip(texts, metadatas)]
- return cls(vectorizer=vectorizer, docs=docs, tfidf_array=tfidf_array, **kwargs)
-
- @classmethod
- def from_documents(
- cls,
- documents: Iterable[Document],
- *,
- tfidf_params: Optional[Dict[str, Any]] = None,
- **kwargs: Any,
- ) -> TFIDFRetriever:
- texts, metadatas = zip(*((d.page_content, d.metadata) for d in documents))
- return cls.from_texts(
- texts=texts, tfidf_params=tfidf_params, metadatas=metadatas, **kwargs
- )
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- from sklearn.metrics.pairwise import cosine_similarity
-
- query_vec = self.vectorizer.transform(
- [query]
- ) # Ip -- (n_docs,x), Op -- (n_docs,n_Feats)
- results = cosine_similarity(self.tfidf_array, query_vec).reshape(
- (-1,)
- ) # Op -- (n_docs,1) -- Cosine Sim with each doc
- return_docs = [self.docs[i] for i in results.argsort()[-self.k :][::-1]]
- return return_docs
-
- def save_local(
- self,
- folder_path: str,
- file_name: str = "tfidf_vectorizer",
- ) -> None:
- try:
- import joblib
- except ImportError:
- raise ImportError(
- "Could not import joblib, please install with `pip install joblib`."
- )
-
- path = Path(folder_path)
- path.mkdir(exist_ok=True, parents=True)
-
- # Save vectorizer with joblib dump.
- joblib.dump(self.vectorizer, path / f"{file_name}.joblib")
-
- # Save docs and tfidf array as pickle.
- with open(path / f"{file_name}.pkl", "wb") as f:
- pickle.dump((self.docs, self.tfidf_array), f)
-
- @classmethod
- def load_local(
- cls,
- folder_path: str,
- *,
- allow_dangerous_deserialization: bool = False,
- file_name: str = "tfidf_vectorizer",
- ) -> TFIDFRetriever:
- """Load the retriever from local storage.
-
- Args:
- folder_path: Folder path to load from.
- allow_dangerous_deserialization: Whether to allow dangerous deserialization.
- Defaults to False.
- The deserialization relies on .joblib and .pkl files, which can be
- modified to deliver a malicious payload that results in execution of
- arbitrary code on your machine. You will need to set this to `True` to
- use deserialization. If you do this, make sure you trust the source of
- the file.
- file_name: File name to load from. Defaults to "tfidf_vectorizer".
-
- Returns:
- TFIDFRetriever: Loaded retriever.
- """
- try:
- import joblib
- except ImportError:
- raise ImportError(
- "Could not import joblib, please install with `pip install joblib`."
- )
-
- if not allow_dangerous_deserialization:
- raise ValueError(
- "The de-serialization of this retriever is based on .joblib and "
- ".pkl files."
- "Such files can be modified to deliver a malicious payload that "
- "results in execution of arbitrary code on your machine."
- "You will need to set `allow_dangerous_deserialization` to `True` to "
- "load this retriever. If you do this, make sure you trust the source "
- "of the file, and you are responsible for validating the file "
- "came from a trusted source."
- )
-
- path = Path(folder_path)
-
- # Load vectorizer with joblib load.
- vectorizer = joblib.load(path / f"{file_name}.joblib")
-
- # Load docs and tfidf array as pickle.
- with open(path / f"{file_name}.pkl", "rb") as f:
- # This code path can only be triggered if the user
- # passed allow_dangerous_deserialization=True
- docs, tfidf_array = pickle.load(f) # ignore[pickle]: explicit-opt-in
-
- return cls(vectorizer=vectorizer, docs=docs, tfidf_array=tfidf_array)
diff --git a/libs/community/langchain_community/retrievers/thirdai_neuraldb.py b/libs/community/langchain_community/retrievers/thirdai_neuraldb.py
deleted file mode 100644
index 2fde6d73d1..0000000000
--- a/libs/community/langchain_community/retrievers/thirdai_neuraldb.py
+++ /dev/null
@@ -1,258 +0,0 @@
-from __future__ import annotations
-
-import importlib
-import os
-from pathlib import Path
-from typing import Any, Dict, List, Optional, Tuple, Union
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init
-from pydantic import ConfigDict, SecretStr
-
-
-class NeuralDBRetriever(BaseRetriever):
- """Document retriever that uses ThirdAI's NeuralDB."""
-
- thirdai_key: SecretStr
- """ThirdAI API Key"""
-
- db: Any = None #: :meta private:
- """NeuralDB instance"""
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- @staticmethod
- def _verify_thirdai_library(thirdai_key: Optional[str] = None) -> None:
- try:
- from thirdai import licensing
-
- importlib.util.find_spec("thirdai.neural_db")
-
- licensing.activate(thirdai_key or os.getenv("THIRDAI_KEY"))
- except ImportError:
- raise ImportError(
- "Could not import thirdai python package and neuraldb dependencies. "
- "Please install it with `pip install thirdai[neural_db]`."
- )
-
- @classmethod
- def from_scratch(
- cls,
- thirdai_key: Optional[str] = None,
- **model_kwargs: dict,
- ) -> NeuralDBRetriever:
- """
- Create a NeuralDBRetriever from scratch.
-
- To use, set the ``THIRDAI_KEY`` environment variable with your ThirdAI
- API key, or pass ``thirdai_key`` as a named parameter.
-
- Example:
- .. code-block:: python
-
- from langchain_community.retrievers import NeuralDBRetriever
-
- retriever = NeuralDBRetriever.from_scratch(
- thirdai_key="your-thirdai-key",
- )
-
- retriever.insert([
- "/path/to/doc.pdf",
- "/path/to/doc.docx",
- "/path/to/doc.csv",
- ])
-
- documents = retriever.invoke("AI-driven music therapy")
- """
- NeuralDBRetriever._verify_thirdai_library(thirdai_key)
- from thirdai import neural_db as ndb
-
- return cls(thirdai_key=thirdai_key, db=ndb.NeuralDB(**model_kwargs)) # type: ignore[arg-type]
-
- @classmethod
- def from_checkpoint(
- cls,
- checkpoint: Union[str, Path],
- thirdai_key: Optional[str] = None,
- ) -> NeuralDBRetriever:
- """
- Create a NeuralDBRetriever with a base model from a saved checkpoint
-
- To use, set the ``THIRDAI_KEY`` environment variable with your ThirdAI
- API key, or pass ``thirdai_key`` as a named parameter.
-
- Example:
- .. code-block:: python
-
- from langchain_community.retrievers import NeuralDBRetriever
-
- retriever = NeuralDBRetriever.from_checkpoint(
- checkpoint="/path/to/checkpoint.ndb",
- thirdai_key="your-thirdai-key",
- )
-
- retriever.insert([
- "/path/to/doc.pdf",
- "/path/to/doc.docx",
- "/path/to/doc.csv",
- ])
-
- documents = retriever.invoke("AI-driven music therapy")
- """
- NeuralDBRetriever._verify_thirdai_library(thirdai_key)
- from thirdai import neural_db as ndb
-
- return cls(thirdai_key=thirdai_key, db=ndb.NeuralDB.from_checkpoint(checkpoint)) # type: ignore[arg-type]
-
- @pre_init
- def validate_environments(cls, values: Dict) -> Dict:
- """Validate ThirdAI environment variables."""
- values["thirdai_key"] = convert_to_secret_str(
- get_from_dict_or_env(
- values,
- "thirdai_key",
- "THIRDAI_KEY",
- )
- )
- return values
-
- def insert(
- self,
- sources: List[Any],
- train: bool = True,
- fast_mode: bool = True,
- **kwargs: dict,
- ) -> None:
- """Inserts files / document sources into the retriever.
-
- Args:
- train: When True this means that the underlying model in the
- NeuralDB will undergo unsupervised pretraining on the inserted files.
- Defaults to True.
- fast_mode: Much faster insertion with a slight drop in performance.
- Defaults to True.
- """
- sources = self._preprocess_sources(sources)
- self.db.insert(
- sources=sources,
- train=train,
- fast_approximation=fast_mode,
- **kwargs,
- )
-
- def _preprocess_sources(self, sources: list) -> list:
- """Checks if the provided sources are string paths. If they are, convert
- to NeuralDB document objects.
-
- Args:
- sources: list of either string paths to PDF, DOCX or CSV files, or
- NeuralDB document objects.
- """
- from thirdai import neural_db as ndb
-
- if not sources:
- return sources
- preprocessed_sources = []
- for doc in sources:
- if not isinstance(doc, str):
- preprocessed_sources.append(doc)
- else:
- if doc.lower().endswith(".pdf"):
- preprocessed_sources.append(ndb.PDF(doc))
- elif doc.lower().endswith(".docx"):
- preprocessed_sources.append(ndb.DOCX(doc))
- elif doc.lower().endswith(".csv"):
- preprocessed_sources.append(ndb.CSV(doc))
- else:
- raise RuntimeError(
- f"Could not automatically load {doc}. Only files "
- "with .pdf, .docx, or .csv extensions can be loaded "
- "automatically. For other formats, please use the "
- "appropriate document object from the ThirdAI library."
- )
- return preprocessed_sources
-
- def upvote(self, query: str, document_id: int) -> None:
- """The retriever upweights the score of a document for a specific query.
- This is useful for fine-tuning the retriever to user behavior.
-
- Args:
- query: text to associate with `document_id`
- document_id: id of the document to associate query with.
- """
- self.db.text_to_result(query, document_id)
-
- def upvote_batch(self, query_id_pairs: List[Tuple[str, int]]) -> None:
- """Given a batch of (query, document id) pairs, the retriever upweights
- the scores of the document for the corresponding queries.
- This is useful for fine-tuning the retriever to user behavior.
-
- Args:
- query_id_pairs: list of (query, document id) pairs. For each pair in
- this list, the model will upweight the document id for the query.
- """
- self.db.text_to_result_batch(query_id_pairs)
-
- def associate(self, source: str, target: str) -> None:
- """The retriever associates a source phrase with a target phrase.
- When the retriever sees the source phrase, it will also consider results
- that are relevant to the target phrase.
-
- Args:
- source: text to associate to `target`.
- target: text to associate `source` to.
- """
- self.db.associate(source, target)
-
- def associate_batch(self, text_pairs: List[Tuple[str, str]]) -> None:
- """Given a batch of (source, target) pairs, the retriever associates
- each source phrase with the corresponding target phrase.
-
- Args:
- text_pairs: list of (source, target) text pairs. For each pair in
- this list, the source will be associated with the target.
- """
- self.db.associate_batch(text_pairs)
-
- def _get_relevant_documents(
- self, query: str, run_manager: CallbackManagerForRetrieverRun, **kwargs: Any
- ) -> List[Document]:
- """Retrieve {top_k} contexts with your retriever for a given query
-
- Args:
- query: Query to submit to the model
- top_k: The max number of context results to retrieve. Defaults to 10.
- """
- try:
- if "top_k" not in kwargs:
- kwargs["top_k"] = 10
- references = self.db.search(query=query, **kwargs)
- return [
- Document(
- page_content=ref.text,
- metadata={
- "id": ref.id,
- "upvote_ids": ref.upvote_ids,
- "source": ref.source,
- "metadata": ref.metadata,
- "score": ref.score,
- "context": ref.context(1),
- },
- )
- for ref in references
- ]
- except Exception as e:
- raise ValueError(f"Error while retrieving documents: {e}") from e
-
- def save(self, path: str) -> None:
- """Saves a NeuralDB instance to disk. Can be loaded into memory by
- calling NeuralDB.from_checkpoint(path)
-
- Args:
- path: path on disk to save the NeuralDB instance to.
- """
- self.db.save(path)
diff --git a/libs/community/langchain_community/retrievers/vespa_retriever.py b/libs/community/langchain_community/retrievers/vespa_retriever.py
deleted file mode 100644
index 6f5eb66aa0..0000000000
--- a/libs/community/langchain_community/retrievers/vespa_retriever.py
+++ /dev/null
@@ -1,126 +0,0 @@
-from __future__ import annotations
-
-import json
-from typing import Any, Dict, List, Literal, Optional, Sequence, Union
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-
-class VespaRetriever(BaseRetriever):
- """`Vespa` retriever."""
-
- app: Any
- """Vespa application to query."""
- body: Dict
- """Body of the query."""
- content_field: str
- """Name of the content field."""
- metadata_fields: Sequence[str]
- """Names of the metadata fields."""
-
- def _query(self, body: Dict) -> List[Document]:
- response = self.app.query(body)
-
- if not str(response.status_code).startswith("2"):
- raise RuntimeError(
- "Could not retrieve data from Vespa. Error code: {}".format(
- response.status_code
- )
- )
-
- root = response.json["root"]
- if "errors" in root:
- raise RuntimeError(json.dumps(root["errors"]))
-
- docs = []
- for child in response.hits:
- page_content = child["fields"].pop(self.content_field, "")
- if self.metadata_fields == "*":
- metadata = child["fields"]
- else:
- metadata = {mf: child["fields"].get(mf) for mf in self.metadata_fields}
- metadata["id"] = child["id"]
- docs.append(Document(page_content=page_content, metadata=metadata))
- return docs
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- body = self.body.copy()
- body["query"] = query
- return self._query(body)
-
- def get_relevant_documents_with_filter(
- self, query: str, *, _filter: Optional[str] = None
- ) -> List[Document]:
- body = self.body.copy()
- _filter = f" and {_filter}" if _filter else ""
- body["yql"] = body["yql"] + _filter
- body["query"] = query
- return self._query(body)
-
- @classmethod
- def from_params(
- cls,
- url: str,
- content_field: str,
- *,
- k: Optional[int] = None,
- metadata_fields: Union[Sequence[str], Literal["*"]] = (),
- sources: Union[Sequence[str], Literal["*"], None] = None,
- _filter: Optional[str] = None,
- yql: Optional[str] = None,
- **kwargs: Any,
- ) -> VespaRetriever:
- """Instantiate retriever from params.
-
- Args:
- url (str): Vespa app URL.
- content_field (str): Field in results to return as Document page_content.
- k (Optional[int]): Number of Documents to return. Defaults to None.
- metadata_fields(Sequence[str] or "*"): Fields in results to include in
- document metadata. Defaults to empty tuple ().
- sources (Sequence[str] or "*" or None): Sources to retrieve
- from. Defaults to None.
- _filter (Optional[str]): Document filter condition expressed in YQL.
- Defaults to None.
- yql (Optional[str]): Full YQL query to be used. Should not be specified
- if _filter or sources are specified. Defaults to None.
- kwargs (Any): Keyword arguments added to query body.
-
- Returns:
- VespaRetriever: Instantiated VespaRetriever.
- """
- try:
- from vespa.application import Vespa
- except ImportError:
- raise ImportError(
- "pyvespa is not installed, please install with `pip install pyvespa`"
- )
- app = Vespa(url)
- body = kwargs.copy()
- if yql and (sources or _filter):
- raise ValueError(
- "yql should only be specified if both sources and _filter are not "
- "specified."
- )
- else:
- if metadata_fields == "*":
- _fields = "*"
- body["summary"] = "short"
- else:
- _fields = ", ".join([content_field] + list(metadata_fields or []))
- _sources = ", ".join(sources) if isinstance(sources, Sequence) else "*"
- _filter = f" and {_filter}" if _filter else ""
- yql = f"select {_fields} from sources {_sources} where userQuery(){_filter}"
- body["yql"] = yql
- if k:
- body["hits"] = k
- return cls(
- app=app,
- body=body,
- content_field=content_field,
- metadata_fields=metadata_fields,
- )
diff --git a/libs/community/langchain_community/retrievers/weaviate_hybrid_search.py b/libs/community/langchain_community/retrievers/weaviate_hybrid_search.py
deleted file mode 100644
index d172d6d6a8..0000000000
--- a/libs/community/langchain_community/retrievers/weaviate_hybrid_search.py
+++ /dev/null
@@ -1,168 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Dict, List, Optional, cast
-from uuid import uuid4
-
-from langchain_core._api import deprecated
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from pydantic import ConfigDict, model_validator
-
-
-@deprecated(
- since="0.3.18",
- removal="1.0",
- alternative_import="langchain_weaviate.WeaviateVectorStore",
-)
-class WeaviateHybridSearchRetriever(BaseRetriever):
- """`Weaviate hybrid search` retriever.
-
- See the documentation:
- https://weaviate.io/blog/hybrid-search-explained
- """
-
- client: Any = None
- """keyword arguments to pass to the Weaviate client."""
- index_name: str
- """The name of the index to use."""
- text_key: str
- """The name of the text key to use."""
- alpha: float = 0.5
- """The weight of the text key in the hybrid search."""
- k: int = 4
- """The number of results to return."""
- attributes: List[str]
- """The attributes to return in the results."""
- create_schema_if_missing: bool = True
- """Whether to create the schema if it doesn't exist."""
-
- @model_validator(mode="before")
- @classmethod
- def validate_client(
- cls,
- values: Dict[str, Any],
- ) -> Any:
- try:
- import weaviate
- except ImportError:
- raise ImportError(
- "Could not import weaviate python package. "
- "Please install it with `pip install weaviate-client`."
- )
- if not isinstance(values["client"], weaviate.Client):
- client = values["client"]
- raise ValueError(
- f"client should be an instance of weaviate.Client, got {type(client)}"
- )
- if values.get("attributes") is None:
- values["attributes"] = []
-
- cast(List, values["attributes"]).append(values["text_key"])
-
- if values.get("create_schema_if_missing", True):
- class_obj = {
- "class": values["index_name"],
- "properties": [{"name": values["text_key"], "dataType": ["text"]}],
- "vectorizer": "text2vec-openai",
- }
-
- if not values["client"].schema.exists(values["index_name"]):
- values["client"].schema.create_class(class_obj)
-
- return values
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- )
-
- # added text_key
- def add_documents(self, docs: List[Document], **kwargs: Any) -> List[str]:
- """Upload documents to Weaviate."""
- from weaviate.util import get_valid_uuid
-
- with self.client.batch as batch:
- ids = []
- for i, doc in enumerate(docs):
- metadata = doc.metadata or {}
- data_properties = {self.text_key: doc.page_content, **metadata}
-
- # If the UUID of one of the objects already exists
- # then the existing objectwill be replaced by the new object.
- if "uuids" in kwargs:
- _id = kwargs["uuids"][i]
- else:
- _id = get_valid_uuid(uuid4())
-
- batch.add_data_object(data_properties, self.index_name, _id)
- ids.append(_id)
- return ids
-
- def _get_relevant_documents(
- self,
- query: str,
- *,
- run_manager: CallbackManagerForRetrieverRun,
- where_filter: Optional[Dict[str, object]] = None,
- score: bool = False,
- hybrid_search_kwargs: Optional[Dict[str, object]] = None,
- ) -> List[Document]:
- """Look up similar documents in Weaviate.
-
- query: The query to search for relevant documents
- of using weviate hybrid search.
-
- where_filter: A filter to apply to the query.
- https://weaviate.io/developers/weaviate/guides/querying/#filtering
-
- score: Whether to include the score, and score explanation
- in the returned Documents meta_data.
-
- hybrid_search_kwargs: Used to pass additional arguments
- to the .with_hybrid() method.
- The primary uses cases for this are:
- 1) Search specific properties only -
- specify which properties to be used during hybrid search portion.
- Note: this is not the same as the (self.attributes) to be returned.
- Example - hybrid_search_kwargs={"properties": ["question", "answer"]}
- https://weaviate.io/developers/weaviate/search/hybrid#selected-properties-only
-
- 2) Weight boosted searched properties -
- Boost the weight of certain properties during the hybrid search portion.
- Example - hybrid_search_kwargs={"properties": ["question^2", "answer"]}
- https://weaviate.io/developers/weaviate/search/hybrid#weight-boost-searched-properties
-
- 3) Search with a custom vector - Define a different vector
- to be used during the hybrid search portion.
- Example - hybrid_search_kwargs={"vector": [0.1, 0.2, 0.3, ...]}
- https://weaviate.io/developers/weaviate/search/hybrid#with-a-custom-vector
-
- 4) Use Fusion ranking method
- Example - from weaviate.gql.get import HybridFusion
- hybrid_search_kwargs={"fusion": fusion_type=HybridFusion.RELATIVE_SCORE}
- https://weaviate.io/developers/weaviate/search/hybrid#fusion-ranking-method
- """
- query_obj = self.client.query.get(self.index_name, self.attributes)
- if where_filter:
- query_obj = query_obj.with_where(where_filter)
-
- if score:
- query_obj = query_obj.with_additional(["score", "explainScore"])
-
- if hybrid_search_kwargs is None:
- hybrid_search_kwargs = {}
-
- result = (
- query_obj.with_hybrid(query, alpha=self.alpha, **hybrid_search_kwargs)
- .with_limit(self.k)
- .do()
- )
- if "errors" in result:
- raise ValueError(f"Error during query: {result['errors']}")
-
- docs = []
-
- for res in result["data"]["Get"][self.index_name]:
- text = res.pop(self.text_key)
- docs.append(Document(page_content=text, metadata=res))
- return docs
diff --git a/libs/community/langchain_community/retrievers/web_research.py b/libs/community/langchain_community/retrievers/web_research.py
deleted file mode 100644
index aa604dd841..0000000000
--- a/libs/community/langchain_community/retrievers/web_research.py
+++ /dev/null
@@ -1,267 +0,0 @@
-import logging
-import re
-from typing import Any, List, Optional
-
-from langchain.chains import LLMChain
-from langchain.chains.prompt_selector import ConditionalPromptSelector
-from langchain_core.callbacks import (
- AsyncCallbackManagerForRetrieverRun,
- CallbackManagerForRetrieverRun,
-)
-from langchain_core.documents import Document
-from langchain_core.language_models import BaseLLM
-from langchain_core.output_parsers import BaseOutputParser
-from langchain_core.prompts import BasePromptTemplate, PromptTemplate
-from langchain_core.retrievers import BaseRetriever
-from langchain_core.vectorstores import VectorStore
-from langchain_text_splitters import RecursiveCharacterTextSplitter, TextSplitter
-from pydantic import BaseModel, Field
-
-from langchain_community.document_loaders import AsyncHtmlLoader
-from langchain_community.document_transformers import Html2TextTransformer
-from langchain_community.llms import LlamaCpp
-from langchain_community.utilities import GoogleSearchAPIWrapper
-
-logger = logging.getLogger(__name__)
-
-
-class SearchQueries(BaseModel):
- """Search queries to research for the user's goal."""
-
- queries: List[str] = Field(
- ..., description="List of search queries to look up on Google"
- )
-
-
-DEFAULT_LLAMA_SEARCH_PROMPT = PromptTemplate(
- input_variables=["question"],
- template="""<> \n You are an assistant tasked with improving Google search \
-results. \n <> \n\n [INST] Generate THREE Google search queries that \
-are similar to this question. The output should be a numbered list of questions \
-and each should have a question mark at the end: \n\n {question} [/INST]""",
-)
-
-DEFAULT_SEARCH_PROMPT = PromptTemplate(
- input_variables=["question"],
- template="""You are an assistant tasked with improving Google search \
-results. Generate THREE Google search queries that are similar to \
-this question. The output should be a numbered list of questions and each \
-should have a question mark at the end: {question}""",
-)
-
-
-class QuestionListOutputParser(BaseOutputParser[List[str]]):
- """Output parser for a list of numbered questions."""
-
- def parse(self, text: str) -> List[str]:
- lines = re.findall(r"\d+\..*?(?:\n|$)", text)
- return lines
-
-
-class WebResearchRetriever(BaseRetriever):
- """`Google Search API` retriever."""
-
- # Inputs
- vectorstore: VectorStore = Field(
- ..., description="Vector store for storing web pages"
- )
- llm_chain: LLMChain
- search: GoogleSearchAPIWrapper = Field(..., description="Google Search API Wrapper")
- num_search_results: int = Field(1, description="Number of pages per Google search")
- text_splitter: TextSplitter = Field(
- RecursiveCharacterTextSplitter(chunk_size=1500, chunk_overlap=50),
- description="Text splitter for splitting web pages into chunks",
- )
- url_database: List[str] = Field(
- default_factory=list, description="List of processed URLs"
- )
- trust_env: bool = Field(
- False,
- description="Whether to use the http_proxy/https_proxy env variables or "
- "check .netrc for proxy configuration",
- )
-
- allow_dangerous_requests: bool = False
- """A flag to force users to acknowledge the risks of SSRF attacks when using
- this retriever.
-
- Users should set this flag to `True` if they have taken the necessary precautions
- to prevent SSRF attacks when using this retriever.
-
- For example, users can run the requests through a properly configured
- proxy and prevent the crawler from accidentally crawling internal resources.
- """
-
- def __init__(self, **kwargs: Any) -> None:
- """Initialize the retriever."""
- allow_dangerous_requests = kwargs.get("allow_dangerous_requests", False)
- if not allow_dangerous_requests:
- raise ValueError(
- "WebResearchRetriever crawls URLs surfaced through "
- "the provided search engine. It is possible that some of those URLs "
- "will end up pointing to machines residing on an internal network, "
- "leading"
- "to an SSRF (Server-Side Request Forgery) attack. "
- "To protect yourself against that risk, you can run the requests "
- "through a proxy and prevent the crawler from accidentally crawling "
- "internal resources."
- "If've taken the necessary precautions, you can set "
- "`allow_dangerous_requests` to `True`."
- )
- super().__init__(**kwargs)
-
- @classmethod
- def from_llm(
- cls,
- vectorstore: VectorStore,
- llm: BaseLLM,
- search: GoogleSearchAPIWrapper,
- prompt: Optional[BasePromptTemplate] = None,
- num_search_results: int = 1,
- text_splitter: RecursiveCharacterTextSplitter = RecursiveCharacterTextSplitter(
- chunk_size=1500, chunk_overlap=150
- ),
- trust_env: bool = False,
- allow_dangerous_requests: bool = False,
- ) -> "WebResearchRetriever":
- """Initialize from llm using default template.
-
- Args:
- vectorstore: Vector store for storing web pages
- llm: llm for search question generation
- search: GoogleSearchAPIWrapper
- prompt: prompt to generating search questions
- num_search_results: Number of pages per Google search
- text_splitter: Text splitter for splitting web pages into chunks
- trust_env: Whether to use the http_proxy/https_proxy env variables
- or check .netrc for proxy configuration
- allow_dangerous_requests: A flag to force users to acknowledge
- the risks of SSRF attacks when using this retriever
-
- Returns:
- WebResearchRetriever
- """
-
- if not prompt:
- QUESTION_PROMPT_SELECTOR = ConditionalPromptSelector(
- default_prompt=DEFAULT_SEARCH_PROMPT,
- conditionals=[
- (lambda llm: isinstance(llm, LlamaCpp), DEFAULT_LLAMA_SEARCH_PROMPT)
- ],
- )
- prompt = QUESTION_PROMPT_SELECTOR.get_prompt(llm)
-
- # Use chat model prompt
- llm_chain = LLMChain(
- llm=llm,
- prompt=prompt,
- output_parser=QuestionListOutputParser(),
- )
-
- return cls(
- vectorstore=vectorstore,
- llm_chain=llm_chain,
- search=search,
- num_search_results=num_search_results,
- text_splitter=text_splitter,
- trust_env=trust_env,
- allow_dangerous_requests=allow_dangerous_requests,
- )
-
- def clean_search_query(self, query: str) -> str:
- # Some search tools (e.g., Google) will
- # fail to return results if query has a
- # leading digit: 1. "LangCh..."
- # Check if the first character is a digit
- if query[0].isdigit():
- # Find the position of the first quote
- first_quote_pos = query.find('"')
- if first_quote_pos != -1:
- # Extract the part of the string after the quote
- query = query[first_quote_pos + 1 :]
- # Remove the trailing quote if present
- if query.endswith('"'):
- query = query[:-1]
- return query.strip()
-
- def search_tool(self, query: str, num_search_results: int = 1) -> List[dict]:
- """Returns num_search_results pages per Google search."""
- query_clean = self.clean_search_query(query)
- result = self.search.results(query_clean, num_search_results)
- return result
-
- def _get_relevant_documents(
- self,
- query: str,
- *,
- run_manager: CallbackManagerForRetrieverRun,
- ) -> List[Document]:
- """Search Google for documents related to the query input.
-
- Args:
- query: user query
-
- Returns:
- Relevant documents from all various urls.
- """
-
- # Get search questions
- logger.info("Generating questions for Google Search ...")
- result = self.llm_chain({"question": query})
- logger.info(f"Questions for Google Search (raw): {result}")
- questions = result["text"]
- logger.info(f"Questions for Google Search: {questions}")
-
- # Get urls
- logger.info("Searching for relevant urls...")
- urls_to_look = []
- for query in questions:
- # Google search
- search_results = self.search_tool(query, self.num_search_results)
- logger.info("Searching for relevant urls...")
- logger.info(f"Search results: {search_results}")
- for res in search_results:
- if res.get("link", None):
- urls_to_look.append(res["link"])
-
- # Relevant urls
- urls = set(urls_to_look)
-
- # Check for any new urls that we have not processed
- new_urls = list(urls.difference(self.url_database))
-
- logger.info(f"New URLs to load: {new_urls}")
- # Load, split, and add new urls to vectorstore
- if new_urls:
- loader = AsyncHtmlLoader(
- new_urls, ignore_load_errors=True, trust_env=self.trust_env
- )
- html2text = Html2TextTransformer()
- logger.info("Indexing new urls...")
- docs = loader.load()
- docs = list(html2text.transform_documents(docs))
- docs = self.text_splitter.split_documents(docs)
- self.vectorstore.add_documents(docs)
- self.url_database.extend(new_urls)
-
- # Search for relevant splits
- # TODO: make this async
- logger.info("Grabbing most relevant splits from urls...")
- docs = []
- for query in questions:
- docs.extend(self.vectorstore.similarity_search(query))
-
- # Get unique docs
- unique_documents_dict = {
- (doc.page_content, tuple(sorted(doc.metadata.items()))): doc for doc in docs
- }
- unique_documents = list(unique_documents_dict.values())
- return unique_documents
-
- async def _aget_relevant_documents(
- self,
- query: str,
- *,
- run_manager: AsyncCallbackManagerForRetrieverRun,
- ) -> List[Document]:
- raise NotImplementedError
diff --git a/libs/community/langchain_community/retrievers/wikipedia.py b/libs/community/langchain_community/retrievers/wikipedia.py
deleted file mode 100644
index 570d7a9aa7..0000000000
--- a/libs/community/langchain_community/retrievers/wikipedia.py
+++ /dev/null
@@ -1,77 +0,0 @@
-from typing import List
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-from langchain_community.utilities.wikipedia import WikipediaAPIWrapper
-
-
-class WikipediaRetriever(BaseRetriever, WikipediaAPIWrapper):
- """`Wikipedia API` retriever.
-
- Setup:
- Install the ``wikipedia`` dependency:
-
- .. code-block:: bash
-
- pip install -U wikipedia
-
- Instantiate:
- .. code-block:: python
-
- from langchain_community.retrievers import WikipediaRetriever
-
- retriever = WikipediaRetriever()
-
- Usage:
- .. code-block:: python
-
- docs = retriever.invoke("TOKYO GHOUL")
- print(docs[0].page_content[:100])
-
- .. code-block:: none
-
- Tokyo Ghoul (Japanese: 東京喰種(トーキョーグール), Hepburn: Tōkyō Gūru) is a Japanese dark fantasy
-
- Use within a chain:
- .. code-block:: python
-
- from langchain_core.output_parsers import StrOutputParser
- from langchain_core.prompts import ChatPromptTemplate
- from langchain_core.runnables import RunnablePassthrough
- from langchain_openai import ChatOpenAI
-
- prompt = ChatPromptTemplate.from_template(
- \"\"\"Answer the question based only on the context provided.
-
- Context: {context}
-
- Question: {question}\"\"\"
- )
-
- llm = ChatOpenAI(model="gpt-3.5-turbo-0125")
-
- def format_docs(docs):
- return "\\n\\n".join(doc.page_content for doc in docs)
-
- chain = (
- {"context": retriever | format_docs, "question": RunnablePassthrough()}
- | prompt
- | llm
- | StrOutputParser()
- )
-
- chain.invoke(
- "Who is the main character in `Tokyo Ghoul` and does he transform into a ghoul?"
- )
-
- .. code-block:: none
-
- 'The main character in Tokyo Ghoul is Ken Kaneki, who transforms into a ghoul after receiving an organ transplant from a ghoul named Rize.'
- """ # noqa: E501
-
- def _get_relevant_documents(
- self, query: str, *, run_manager: CallbackManagerForRetrieverRun
- ) -> List[Document]:
- return self.load(query=query)
diff --git a/libs/community/langchain_community/retrievers/you.py b/libs/community/langchain_community/retrievers/you.py
deleted file mode 100644
index 8ce080545a..0000000000
--- a/libs/community/langchain_community/retrievers/you.py
+++ /dev/null
@@ -1,39 +0,0 @@
-from typing import Any, List
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForRetrieverRun,
- CallbackManagerForRetrieverRun,
-)
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-
-from langchain_community.utilities import YouSearchAPIWrapper
-
-
-class YouRetriever(BaseRetriever, YouSearchAPIWrapper):
- """You.com Search API retriever.
-
- It wraps results() to get_relevant_documents
- It uses all YouSearchAPIWrapper arguments without any change.
- """
-
- def _get_relevant_documents(
- self,
- query: str,
- *,
- run_manager: CallbackManagerForRetrieverRun,
- **kwargs: Any,
- ) -> List[Document]:
- return self.results(query, run_manager=run_manager.get_child(), **kwargs)
-
- async def _aget_relevant_documents(
- self,
- query: str,
- *,
- run_manager: AsyncCallbackManagerForRetrieverRun,
- **kwargs: Any,
- ) -> List[Document]:
- results = await self.results_async(
- query, run_manager=run_manager.get_child(), **kwargs
- )
- return results
diff --git a/libs/community/langchain_community/retrievers/zep.py b/libs/community/langchain_community/retrievers/zep.py
deleted file mode 100644
index d59aa00781..0000000000
--- a/libs/community/langchain_community/retrievers/zep.py
+++ /dev/null
@@ -1,183 +0,0 @@
-from __future__ import annotations
-
-from enum import Enum
-from typing import TYPE_CHECKING, Any, Dict, List, Optional
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForRetrieverRun,
- CallbackManagerForRetrieverRun,
-)
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from pydantic import model_validator
-
-if TYPE_CHECKING:
- from zep_python.memory import MemorySearchResult
-
-
-class SearchScope(str, Enum):
- """Which documents to search. Messages or Summaries?"""
-
- messages = "messages"
- """Search chat history messages."""
- summary = "summary"
- """Search chat history summaries."""
-
-
-class SearchType(str, Enum):
- """Enumerator of the types of search to perform."""
-
- similarity = "similarity"
- """Similarity search."""
- mmr = "mmr"
- """Maximal Marginal Relevance reranking of similarity search."""
-
-
-class ZepRetriever(BaseRetriever):
- """`Zep` MemoryStore Retriever.
-
- Search your user's long-term chat history with Zep.
-
- Zep offers both simple semantic search and Maximal Marginal Relevance (MMR)
- reranking of search results.
-
- Note: You will need to provide the user's `session_id` to use this retriever.
-
- Args:
- url: URL of your Zep server (required)
- api_key: Your Zep API key (optional)
- session_id: Identifies your user or a user's session (required)
- top_k: Number of documents to return (default: 3, optional)
- search_type: Type of search to perform (similarity / mmr) (default: similarity,
- optional)
- mmr_lambda: Lambda value for MMR search. Defaults to 0.5 (optional)
-
- Zep - Fast, scalable building blocks for LLM Apps
- =========
- Zep is an open source platform for productionizing LLM apps. Go from a prototype
- built in LangChain or LlamaIndex, or a custom app, to production in minutes without
- rewriting code.
-
- For server installation instructions, see:
- https://docs.getzep.com/deployment/quickstart/
- """
-
- zep_client: Optional[Any] = None
- """Zep client."""
- url: str
- """URL of your Zep server."""
- api_key: Optional[str] = None
- """Your Zep API key."""
- session_id: str
- """Zep session ID."""
- top_k: Optional[int]
- """Number of items to return."""
- search_scope: SearchScope = SearchScope.messages
- """Which documents to search. Messages or Summaries?"""
- search_type: SearchType = SearchType.similarity
- """Type of search to perform (similarity / mmr)"""
- mmr_lambda: Optional[float] = None
- """Lambda value for MMR search."""
-
- @model_validator(mode="before")
- @classmethod
- def create_client(cls, values: dict) -> Any:
- try:
- from zep_python import ZepClient
- except ImportError:
- raise ImportError(
- "Could not import zep-python package. "
- "Please install it with `pip install zep-python`."
- )
- values["zep_client"] = values.get(
- "zep_client",
- ZepClient(base_url=values["url"], api_key=values.get("api_key")),
- )
- return values
-
- def _messages_search_result_to_doc(
- self, results: List[MemorySearchResult]
- ) -> List[Document]:
- return [
- Document(
- page_content=r.message.pop("content"),
- metadata={"score": r.dist, **r.message},
- )
- for r in results
- if r.message
- ]
-
- def _summary_search_result_to_doc(
- self, results: List[MemorySearchResult]
- ) -> List[Document]:
- return [
- Document(
- page_content=r.summary.content,
- metadata={
- "score": r.dist,
- "uuid": r.summary.uuid,
- "created_at": r.summary.created_at,
- "token_count": r.summary.token_count,
- },
- )
- for r in results
- if r.summary
- ]
-
- def _get_relevant_documents(
- self,
- query: str,
- *,
- run_manager: CallbackManagerForRetrieverRun,
- metadata: Optional[Dict[str, Any]] = None,
- ) -> List[Document]:
- from zep_python.memory import MemorySearchPayload
-
- if not self.zep_client:
- raise RuntimeError("Zep client not initialized.")
-
- payload = MemorySearchPayload(
- text=query,
- metadata=metadata,
- search_scope=self.search_scope,
- search_type=self.search_type,
- mmr_lambda=self.mmr_lambda,
- )
-
- results: List[MemorySearchResult] = self.zep_client.memory.search_memory(
- self.session_id, payload, limit=self.top_k
- )
-
- if self.search_scope == SearchScope.summary:
- return self._summary_search_result_to_doc(results)
-
- return self._messages_search_result_to_doc(results)
-
- async def _aget_relevant_documents(
- self,
- query: str,
- *,
- run_manager: AsyncCallbackManagerForRetrieverRun,
- metadata: Optional[Dict[str, Any]] = None,
- ) -> List[Document]:
- from zep_python.memory import MemorySearchPayload
-
- if not self.zep_client:
- raise RuntimeError("Zep client not initialized.")
-
- payload = MemorySearchPayload(
- text=query,
- metadata=metadata,
- search_scope=self.search_scope,
- search_type=self.search_type,
- mmr_lambda=self.mmr_lambda,
- )
-
- results: List[MemorySearchResult] = await self.zep_client.memory.asearch_memory(
- self.session_id, payload, limit=self.top_k
- )
-
- if self.search_scope == SearchScope.summary:
- return self._summary_search_result_to_doc(results)
-
- return self._messages_search_result_to_doc(results)
diff --git a/libs/community/langchain_community/retrievers/zep_cloud.py b/libs/community/langchain_community/retrievers/zep_cloud.py
deleted file mode 100644
index c4e3f11040..0000000000
--- a/libs/community/langchain_community/retrievers/zep_cloud.py
+++ /dev/null
@@ -1,163 +0,0 @@
-from __future__ import annotations
-
-from typing import TYPE_CHECKING, Any, Dict, List, Optional
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForRetrieverRun,
- CallbackManagerForRetrieverRun,
-)
-from langchain_core.documents import Document
-from langchain_core.retrievers import BaseRetriever
-from pydantic import model_validator
-
-if TYPE_CHECKING:
- from zep_cloud import MemorySearchResult, SearchScope, SearchType
- from zep_cloud.client import AsyncZep, Zep
-
-
-class ZepCloudRetriever(BaseRetriever):
- """`Zep Cloud` MemoryStore Retriever.
-
- Search your user's long-term chat history with Zep.
-
- Zep offers both simple semantic search and Maximal Marginal Relevance (MMR)
- reranking of search results.
-
- Note: You will need to provide the user's `session_id` to use this retriever.
-
- Args:
- api_key: Your Zep API key
- session_id: Identifies your user or a user's session (required)
- top_k: Number of documents to return (default: 3, optional)
- search_type: Type of search to perform (similarity / mmr)
- (default: similarity, optional)
- mmr_lambda: Lambda value for MMR search. Defaults to 0.5 (optional)
-
- Zep - Recall, understand, and extract data from chat histories.
- Power personalized AI experiences.
- =========
- Zep is a long-term memory service for AI Assistant apps.
- With Zep, you can provide AI assistants with the ability
- to recall past conversations,
- no matter how distant, while also reducing hallucinations, latency, and cost.
-
- see Zep Cloud Docs: https://help.getzep.com
- """
-
- api_key: str
- """Your Zep API key."""
- zep_client: Zep
- """Zep client used for making API requests."""
- zep_client_async: AsyncZep
- """Async Zep client used for making API requests."""
- session_id: str
- """Zep session ID."""
- top_k: Optional[int]
- """Number of items to return."""
- search_scope: SearchScope = "messages"
- """Which documents to search. Messages or Summaries?"""
- search_type: SearchType = "similarity"
- """Type of search to perform (similarity / mmr)"""
- mmr_lambda: Optional[float] = None
- """Lambda value for MMR search."""
-
- @model_validator(mode="before")
- @classmethod
- def create_client(cls, values: dict) -> Any:
- try:
- from zep_cloud.client import AsyncZep, Zep
- except ImportError:
- raise ImportError(
- "Could not import zep-cloud package. "
- "Please install it with `pip install zep-cloud`."
- )
- if values.get("api_key") is None:
- raise ValueError("Zep API key is required.")
- values["zep_client"] = Zep(api_key=values.get("api_key"))
- values["zep_client_async"] = AsyncZep(api_key=values.get("api_key"))
- return values
-
- def _messages_search_result_to_doc(
- self, results: List[MemorySearchResult]
- ) -> List[Document]:
- return [
- Document(
- page_content=str(r.message.content),
- metadata={
- "score": r.score,
- "uuid": r.message.uuid_,
- "created_at": r.message.created_at,
- "token_count": r.message.token_count,
- "role": r.message.role or r.message.role_type,
- },
- )
- for r in results or []
- if r.message
- ]
-
- def _summary_search_result_to_doc(
- self, results: List[MemorySearchResult]
- ) -> List[Document]:
- return [
- Document(
- page_content=str(r.summary.content),
- metadata={
- "score": r.score,
- "uuid": r.summary.uuid_,
- "created_at": r.summary.created_at,
- "token_count": r.summary.token_count,
- },
- )
- for r in results
- if r.summary
- ]
-
- def _get_relevant_documents(
- self,
- query: str,
- *,
- run_manager: CallbackManagerForRetrieverRun,
- metadata: Optional[Dict[str, Any]] = None,
- ) -> List[Document]:
- if not self.zep_client:
- raise RuntimeError("Zep client not initialized.")
-
- results = self.zep_client.memory.search(
- self.session_id,
- text=query,
- metadata=metadata,
- search_scope=self.search_scope,
- search_type=self.search_type,
- mmr_lambda=self.mmr_lambda,
- limit=self.top_k,
- )
-
- if self.search_scope == "summary":
- return self._summary_search_result_to_doc(results)
-
- return self._messages_search_result_to_doc(results)
-
- async def _aget_relevant_documents(
- self,
- query: str,
- *,
- run_manager: AsyncCallbackManagerForRetrieverRun,
- metadata: Optional[Dict[str, Any]] = None,
- ) -> List[Document]:
- if not self.zep_client_async:
- raise RuntimeError("Zep client not initialized.")
-
- results = await self.zep_client_async.memory.search(
- self.session_id,
- text=query,
- metadata=metadata,
- search_scope=self.search_scope,
- search_type=self.search_type,
- mmr_lambda=self.mmr_lambda,
- limit=self.top_k,
- )
-
- if self.search_scope == "summary":
- return self._summary_search_result_to_doc(results)
-
- return self._messages_search_result_to_doc(results)
diff --git a/libs/community/langchain_community/retrievers/zilliz.py b/libs/community/langchain_community/retrievers/zilliz.py
deleted file mode 100644
index d7e6149420..0000000000
--- a/libs/community/langchain_community/retrievers/zilliz.py
+++ /dev/null
@@ -1,87 +0,0 @@
-import warnings
-from typing import Any, Dict, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForRetrieverRun
-from langchain_core.documents import Document
-from langchain_core.embeddings import Embeddings
-from langchain_core.retrievers import BaseRetriever
-from pydantic import model_validator
-
-from langchain_community.vectorstores.zilliz import Zilliz
-
-# TODO: Update to ZillizClient + Hybrid Search when available
-
-
-class ZillizRetriever(BaseRetriever):
- """`Zilliz API` retriever."""
-
- embedding_function: Embeddings
- """The underlying embedding function from which documents will be retrieved."""
- collection_name: str = "LangChainCollection"
- """The name of the collection in Zilliz."""
- connection_args: Optional[Dict[str, Any]] = None
- """The connection arguments for the Zilliz client."""
- consistency_level: str = "Session"
- """The consistency level for the Zilliz client."""
- search_params: Optional[dict] = None
- """The search parameters for the Zilliz client."""
- store: Zilliz
- """The underlying Zilliz store."""
- retriever: BaseRetriever
- """The underlying retriever."""
-
- @model_validator(mode="before")
- @classmethod
- def create_client(cls, values: dict) -> Any:
- values["store"] = Zilliz(
- values["embedding_function"],
- values["collection_name"],
- values["connection_args"],
- values["consistency_level"],
- )
- values["retriever"] = values["store"].as_retriever(
- search_kwargs={"param": values["search_params"]}
- )
- return values
-
- def add_texts(
- self, texts: List[str], metadatas: Optional[List[dict]] = None
- ) -> None:
- """Add text to the Zilliz store
-
- Args:
- texts (List[str]): The text
- metadatas (List[dict]): Metadata dicts, must line up with existing store
- """
- self.store.add_texts(texts, metadatas)
-
- def _get_relevant_documents(
- self,
- query: str,
- *,
- run_manager: CallbackManagerForRetrieverRun,
- **kwargs: Any,
- ) -> List[Document]:
- return self.retriever.invoke(
- query, run_manager=run_manager.get_child(), **kwargs
- )
-
-
-def ZillizRetreiver(*args: Any, **kwargs: Any) -> ZillizRetriever:
- """Deprecated ZillizRetreiver.
-
- Please use ZillizRetriever ('i' before 'e') instead.
-
- Args:
- *args:
- **kwargs:
-
- Returns:
- ZillizRetriever
- """
- warnings.warn(
- "ZillizRetreiver will be deprecated in the future. "
- "Please use ZillizRetriever ('i' before 'e') instead.",
- DeprecationWarning,
- )
- return ZillizRetriever(*args, **kwargs)
diff --git a/libs/community/langchain_community/storage/__init__.py b/libs/community/langchain_community/storage/__init__.py
deleted file mode 100644
index 21a6090bd1..0000000000
--- a/libs/community/langchain_community/storage/__init__.py
+++ /dev/null
@@ -1,69 +0,0 @@
-"""**Storage** is an implementation of key-value store.
-
-Storage module provides implementations of various key-value stores that conform
-to a simple key-value interface.
-
-The primary goal of these storages is to support caching.
-
-
-**Class hierarchy:**
-
-.. code-block::
-
- BaseStore --> Store # Examples: MongoDBStore, RedisStore
-
-"""
-
-import importlib
-from typing import TYPE_CHECKING, Any
-
-if TYPE_CHECKING:
- from langchain_community.storage.astradb import (
- AstraDBByteStore,
- AstraDBStore,
- )
- from langchain_community.storage.cassandra import (
- CassandraByteStore,
- )
- from langchain_community.storage.mongodb import MongoDBByteStore, MongoDBStore
- from langchain_community.storage.redis import (
- RedisStore,
- )
- from langchain_community.storage.sql import (
- SQLStore,
- )
- from langchain_community.storage.upstash_redis import (
- UpstashRedisByteStore,
- UpstashRedisStore,
- )
-
-__all__ = [
- "AstraDBByteStore",
- "AstraDBStore",
- "CassandraByteStore",
- "MongoDBStore",
- "MongoDBByteStore",
- "RedisStore",
- "SQLStore",
- "UpstashRedisByteStore",
- "UpstashRedisStore",
-]
-
-_module_lookup = {
- "AstraDBByteStore": "langchain_community.storage.astradb",
- "AstraDBStore": "langchain_community.storage.astradb",
- "CassandraByteStore": "langchain_community.storage.cassandra",
- "MongoDBStore": "langchain_community.storage.mongodb",
- "MongoDBByteStore": "langchain_community.storage.mongodb",
- "RedisStore": "langchain_community.storage.redis",
- "SQLStore": "langchain_community.storage.sql",
- "UpstashRedisByteStore": "langchain_community.storage.upstash_redis",
- "UpstashRedisStore": "langchain_community.storage.upstash_redis",
-}
-
-
-def __getattr__(name: str) -> Any:
- if name in _module_lookup:
- module = importlib.import_module(_module_lookup[name])
- return getattr(module, name)
- raise AttributeError(f"module {__name__} has no attribute {name}")
diff --git a/libs/community/langchain_community/storage/astradb.py b/libs/community/langchain_community/storage/astradb.py
deleted file mode 100644
index be6a6a32d0..0000000000
--- a/libs/community/langchain_community/storage/astradb.py
+++ /dev/null
@@ -1,238 +0,0 @@
-from __future__ import annotations
-
-import base64
-from abc import ABC, abstractmethod
-from typing import (
- TYPE_CHECKING,
- Any,
- AsyncIterator,
- Generic,
- Iterator,
- List,
- Optional,
- Sequence,
- Tuple,
- TypeVar,
-)
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.stores import BaseStore, ByteStore
-
-from langchain_community.utilities.astradb import (
- SetupMode,
- _AstraDBCollectionEnvironment,
-)
-
-if TYPE_CHECKING:
- from astrapy.db import AstraDB, AsyncAstraDB
-
-V = TypeVar("V")
-
-
-class AstraDBBaseStore(Generic[V], BaseStore[str, V], ABC):
- """Base class for the DataStax AstraDB data store."""
-
- def __init__(self, *args: Any, **kwargs: Any) -> None:
- self.astra_env = _AstraDBCollectionEnvironment(*args, **kwargs)
- self.collection = self.astra_env.collection
- self.async_collection = self.astra_env.async_collection
-
- @abstractmethod
- def decode_value(self, value: Any) -> Optional[V]:
- """Decodes value from Astra DB"""
-
- @abstractmethod
- def encode_value(self, value: Optional[V]) -> Any:
- """Encodes value for Astra DB"""
-
- def mget(self, keys: Sequence[str]) -> List[Optional[V]]:
- self.astra_env.ensure_db_setup()
- docs_dict = {}
- for doc in self.collection.paginated_find(filter={"_id": {"$in": list(keys)}}):
- docs_dict[doc["_id"]] = doc.get("value")
- return [self.decode_value(docs_dict.get(key)) for key in keys]
-
- async def amget(self, keys: Sequence[str]) -> List[Optional[V]]:
- await self.astra_env.aensure_db_setup()
- docs_dict = {}
- async for doc in self.async_collection.paginated_find(
- filter={"_id": {"$in": list(keys)}}
- ):
- docs_dict[doc["_id"]] = doc.get("value")
- return [self.decode_value(docs_dict.get(key)) for key in keys]
-
- def mset(self, key_value_pairs: Sequence[Tuple[str, V]]) -> None:
- self.astra_env.ensure_db_setup()
- for k, v in key_value_pairs:
- self.collection.upsert({"_id": k, "value": self.encode_value(v)})
-
- async def amset(self, key_value_pairs: Sequence[Tuple[str, V]]) -> None:
- await self.astra_env.aensure_db_setup()
- for k, v in key_value_pairs:
- await self.async_collection.upsert(
- {
- "_id": k,
- "value": self.encode_value(v),
- }
- )
-
- def mdelete(self, keys: Sequence[str]) -> None:
- self.astra_env.ensure_db_setup()
- self.collection.delete_many(filter={"_id": {"$in": list(keys)}})
-
- async def amdelete(self, keys: Sequence[str]) -> None:
- await self.astra_env.aensure_db_setup()
- await self.async_collection.delete_many(filter={"_id": {"$in": list(keys)}})
-
- def yield_keys(self, *, prefix: Optional[str] = None) -> Iterator[str]:
- self.astra_env.ensure_db_setup()
- docs = self.collection.paginated_find()
- for doc in docs:
- key = doc["_id"]
- if not prefix or key.startswith(prefix):
- yield key
-
- async def ayield_keys(self, *, prefix: Optional[str] = None) -> AsyncIterator[str]:
- await self.astra_env.aensure_db_setup()
- async for doc in self.async_collection.paginated_find():
- key = doc["_id"]
- if not prefix or key.startswith(prefix):
- yield key
-
-
-@deprecated(
- since="0.0.22",
- removal="1.0",
- alternative_import="langchain_astradb.AstraDBStore",
-)
-class AstraDBStore(AstraDBBaseStore[Any]):
- def __init__(
- self,
- collection_name: str,
- token: Optional[str] = None,
- api_endpoint: Optional[str] = None,
- astra_db_client: Optional[AstraDB] = None,
- namespace: Optional[str] = None,
- *,
- async_astra_db_client: Optional[AsyncAstraDB] = None,
- pre_delete_collection: bool = False,
- setup_mode: SetupMode = SetupMode.SYNC,
- ) -> None:
- """BaseStore implementation using DataStax AstraDB as the underlying store.
-
- The value type can be any type serializable by json.dumps.
- Can be used to store embeddings with the CacheBackedEmbeddings.
-
- Documents in the AstraDB collection will have the format
-
- .. code-block:: json
-
- {
- "_id": "",
- "value":
- }
-
- Args:
- collection_name: name of the Astra DB collection to create/use.
- token: API token for Astra DB usage.
- api_endpoint: full URL to the API endpoint,
- such as `https://-us-east1.apps.astra.datastax.com`.
- astra_db_client: *alternative to token+api_endpoint*,
- you can pass an already-created 'astrapy.db.AstraDB' instance.
- async_astra_db_client: *alternative to token+api_endpoint*,
- you can pass an already-created 'astrapy.db.AsyncAstraDB' instance.
- namespace: namespace (aka keyspace) where the
- collection is created. Defaults to the database's "default namespace".
- setup_mode: mode used to create the Astra DB collection (SYNC, ASYNC or
- OFF).
- pre_delete_collection: whether to delete the collection
- before creating it. If False and the collection already exists,
- the collection will be used as is.
- """
- # Constructor doc is not inherited so we have to override it.
- super().__init__(
- collection_name=collection_name,
- token=token,
- api_endpoint=api_endpoint,
- astra_db_client=astra_db_client,
- async_astra_db_client=async_astra_db_client,
- namespace=namespace,
- setup_mode=setup_mode,
- pre_delete_collection=pre_delete_collection,
- )
-
- def decode_value(self, value: Any) -> Any:
- return value
-
- def encode_value(self, value: Any) -> Any:
- return value
-
-
-@deprecated(
- since="0.0.22",
- removal="1.0",
- alternative_import="langchain_astradb.AstraDBByteStore",
-)
-class AstraDBByteStore(AstraDBBaseStore[bytes], ByteStore):
- def __init__(
- self,
- collection_name: str,
- token: Optional[str] = None,
- api_endpoint: Optional[str] = None,
- astra_db_client: Optional[AstraDB] = None,
- namespace: Optional[str] = None,
- *,
- async_astra_db_client: Optional[AsyncAstraDB] = None,
- pre_delete_collection: bool = False,
- setup_mode: SetupMode = SetupMode.SYNC,
- ) -> None:
- """ByteStore implementation using DataStax AstraDB as the underlying store.
-
- The bytes values are converted to base64 encoded strings
- Documents in the AstraDB collection will have the format
-
- .. code-block:: json
-
- {
- "_id": "",
- "value": ""
- }
-
- Args:
- collection_name: name of the Astra DB collection to create/use.
- token: API token for Astra DB usage.
- api_endpoint: full URL to the API endpoint,
- such as `https://-us-east1.apps.astra.datastax.com`.
- astra_db_client: *alternative to token+api_endpoint*,
- you can pass an already-created 'astrapy.db.AstraDB' instance.
- async_astra_db_client: *alternative to token+api_endpoint*,
- you can pass an already-created 'astrapy.db.AsyncAstraDB' instance.
- namespace: namespace (aka keyspace) where the
- collection is created. Defaults to the database's "default namespace".
- setup_mode: mode used to create the Astra DB collection (SYNC, ASYNC or
- OFF).
- pre_delete_collection: whether to delete the collection
- before creating it. If False and the collection already exists,
- the collection will be used as is.
- """
- # Constructor doc is not inherited so we have to override it.
- super().__init__(
- collection_name=collection_name,
- token=token,
- api_endpoint=api_endpoint,
- astra_db_client=astra_db_client,
- async_astra_db_client=async_astra_db_client,
- namespace=namespace,
- setup_mode=setup_mode,
- pre_delete_collection=pre_delete_collection,
- )
-
- def decode_value(self, value: Any) -> Optional[bytes]:
- if value is None:
- return None
- return base64.b64decode(value)
-
- def encode_value(self, value: Optional[bytes]) -> Any:
- if value is None:
- return None
- return base64.b64encode(value).decode("ascii")
diff --git a/libs/community/langchain_community/storage/cassandra.py b/libs/community/langchain_community/storage/cassandra.py
deleted file mode 100644
index d2d97a3557..0000000000
--- a/libs/community/langchain_community/storage/cassandra.py
+++ /dev/null
@@ -1,220 +0,0 @@
-from __future__ import annotations
-
-import asyncio
-from asyncio import InvalidStateError, Task
-from typing import (
- TYPE_CHECKING,
- AsyncIterator,
- Iterator,
- List,
- Optional,
- Sequence,
- Tuple,
-)
-
-from langchain_core.stores import ByteStore
-
-from langchain_community.utilities.cassandra import SetupMode, aexecute_cql
-
-if TYPE_CHECKING:
- from cassandra.cluster import Session
- from cassandra.query import PreparedStatement
-
-CREATE_TABLE_CQL_TEMPLATE = """
- CREATE TABLE IF NOT EXISTS {keyspace}.{table}
- (row_id TEXT, body_blob BLOB, PRIMARY KEY (row_id));
-"""
-SELECT_TABLE_CQL_TEMPLATE = (
- """SELECT row_id, body_blob FROM {keyspace}.{table} WHERE row_id IN ?;"""
-)
-SELECT_ALL_TABLE_CQL_TEMPLATE = """SELECT row_id, body_blob FROM {keyspace}.{table};"""
-INSERT_TABLE_CQL_TEMPLATE = (
- """INSERT INTO {keyspace}.{table} (row_id, body_blob) VALUES (?, ?);"""
-)
-DELETE_TABLE_CQL_TEMPLATE = """DELETE FROM {keyspace}.{table} WHERE row_id IN ?;"""
-
-
-class CassandraByteStore(ByteStore):
- """A ByteStore implementation using Cassandra as the backend.
-
- Parameters:
- table: The name of the table to use.
- session: A Cassandra session object. If not provided, it will be resolved
- from the cassio config.
- keyspace: The keyspace to use. If not provided, it will be resolved
- from the cassio config.
- setup_mode: The setup mode to use. Default is SYNC (SetupMode.SYNC).
- """
-
- def __init__(
- self,
- table: str,
- *,
- session: Optional[Session] = None,
- keyspace: Optional[str] = None,
- setup_mode: SetupMode = SetupMode.SYNC,
- ) -> None:
- if not session or not keyspace:
- try:
- from cassio.config import check_resolve_keyspace, check_resolve_session
-
- self.keyspace = keyspace or check_resolve_keyspace(keyspace)
- self.session = session or check_resolve_session()
- except (ImportError, ModuleNotFoundError):
- raise ImportError(
- "Could not import a recent cassio package."
- "Please install it with `pip install --upgrade cassio`."
- )
- else:
- self.keyspace = keyspace
- self.session = session
- self.table = table
- self.select_statement = None
- self.insert_statement = None
- self.delete_statement = None
-
- create_cql = CREATE_TABLE_CQL_TEMPLATE.format(
- keyspace=self.keyspace,
- table=self.table,
- )
- self.db_setup_task: Optional[Task[None]] = None
- if setup_mode == SetupMode.ASYNC:
- self.db_setup_task = asyncio.create_task(
- aexecute_cql(self.session, create_cql)
- )
- else:
- self.session.execute(create_cql)
-
- def ensure_db_setup(self) -> None:
- """Ensure that the DB setup is finished. If not, raise a ValueError."""
- if self.db_setup_task:
- try:
- self.db_setup_task.result()
- except InvalidStateError:
- raise ValueError(
- "Asynchronous setup of the DB not finished. "
- "NB: AstraDB components sync methods shouldn't be called from the "
- "event loop. Consider using their async equivalents."
- )
-
- async def aensure_db_setup(self) -> None:
- """Ensure that the DB setup is finished. If not, wait for it."""
- if self.db_setup_task:
- await self.db_setup_task
-
- def get_select_statement(self) -> PreparedStatement:
- """Get the prepared select statement for the table.
- If not available, prepare it.
-
- Returns:
- PreparedStatement: The prepared statement.
- """
- if not self.select_statement:
- self.select_statement = self.session.prepare(
- SELECT_TABLE_CQL_TEMPLATE.format(
- keyspace=self.keyspace, table=self.table
- )
- )
- return self.select_statement
-
- def get_insert_statement(self) -> PreparedStatement:
- """Get the prepared insert statement for the table.
- If not available, prepare it.
-
- Returns:
- PreparedStatement: The prepared statement.
- """
- if not self.insert_statement:
- self.insert_statement = self.session.prepare(
- INSERT_TABLE_CQL_TEMPLATE.format(
- keyspace=self.keyspace, table=self.table
- )
- )
- return self.insert_statement
-
- def get_delete_statement(self) -> PreparedStatement:
- """Get the prepared delete statement for the table.
- If not available, prepare it.
-
- Returns:
- PreparedStatement: The prepared statement.
- """
-
- if not self.delete_statement:
- self.delete_statement = self.session.prepare(
- DELETE_TABLE_CQL_TEMPLATE.format(
- keyspace=self.keyspace, table=self.table
- )
- )
- return self.delete_statement
-
- def mget(self, keys: Sequence[str]) -> List[Optional[bytes]]:
- from cassandra.query import ValueSequence
-
- self.ensure_db_setup()
- docs_dict = {}
- for row in self.session.execute(
- self.get_select_statement(), [ValueSequence(keys)]
- ):
- docs_dict[row.row_id] = row.body_blob
- return [docs_dict.get(key) for key in keys]
-
- async def amget(self, keys: Sequence[str]) -> List[Optional[bytes]]:
- from cassandra.query import ValueSequence
-
- await self.aensure_db_setup()
- docs_dict = {}
- for row in await aexecute_cql(
- self.session, self.get_select_statement(), parameters=[ValueSequence(keys)]
- ):
- docs_dict[row.row_id] = row.body_blob
- return [docs_dict.get(key) for key in keys]
-
- def mset(self, key_value_pairs: Sequence[Tuple[str, bytes]]) -> None:
- self.ensure_db_setup()
- insert_statement = self.get_insert_statement()
- for k, v in key_value_pairs:
- self.session.execute(insert_statement, (k, v))
-
- async def amset(self, key_value_pairs: Sequence[Tuple[str, bytes]]) -> None:
- await self.aensure_db_setup()
- insert_statement = self.get_insert_statement()
- for k, v in key_value_pairs:
- await aexecute_cql(self.session, insert_statement, parameters=(k, v))
-
- def mdelete(self, keys: Sequence[str]) -> None:
- from cassandra.query import ValueSequence
-
- self.ensure_db_setup()
- self.session.execute(self.get_delete_statement(), [ValueSequence(keys)])
-
- async def amdelete(self, keys: Sequence[str]) -> None:
- from cassandra.query import ValueSequence
-
- await self.aensure_db_setup()
- await aexecute_cql(
- self.session, self.get_delete_statement(), parameters=[ValueSequence(keys)]
- )
-
- def yield_keys(self, *, prefix: Optional[str] = None) -> Iterator[str]:
- self.ensure_db_setup()
- for row in self.session.execute(
- SELECT_ALL_TABLE_CQL_TEMPLATE.format(
- keyspace=self.keyspace, table=self.table
- )
- ):
- key = row.row_id
- if not prefix or key.startswith(prefix):
- yield key
-
- async def ayield_keys(self, *, prefix: Optional[str] = None) -> AsyncIterator[str]:
- await self.aensure_db_setup()
- for row in await aexecute_cql(
- self.session,
- SELECT_ALL_TABLE_CQL_TEMPLATE.format(
- keyspace=self.keyspace, table=self.table
- ),
- ):
- key = row.row_id
- if not prefix or key.startswith(prefix):
- yield key
diff --git a/libs/community/langchain_community/storage/exceptions.py b/libs/community/langchain_community/storage/exceptions.py
deleted file mode 100644
index 82d7c8a2fa..0000000000
--- a/libs/community/langchain_community/storage/exceptions.py
+++ /dev/null
@@ -1,3 +0,0 @@
-from langchain_core.stores import InvalidKeyException
-
-__all__ = ["InvalidKeyException"]
diff --git a/libs/community/langchain_community/storage/mongodb.py b/libs/community/langchain_community/storage/mongodb.py
deleted file mode 100644
index 264faefcba..0000000000
--- a/libs/community/langchain_community/storage/mongodb.py
+++ /dev/null
@@ -1,248 +0,0 @@
-from typing import Iterator, List, Optional, Sequence, Tuple
-
-from langchain_core.documents import Document
-from langchain_core.stores import BaseStore
-
-
-class MongoDBByteStore(BaseStore[str, bytes]):
- """BaseStore implementation using MongoDB as the underlying store.
-
- Examples:
- Create a MongoDBByteStore instance and perform operations on it:
-
- .. code-block:: python
-
- # Instantiate the MongoDBByteStore with a MongoDB connection
- from langchain.storage import MongoDBByteStore
-
- mongo_conn_str = "mongodb://localhost:27017/"
- mongodb_store = MongoDBBytesStore(mongo_conn_str, db_name="test-db",
- collection_name="test-collection")
-
- # Set values for keys
- mongodb_store.mset([("key1", "hello"), ("key2", "workd")])
-
- # Get values for keys
- values = mongodb_store.mget(["key1", "key2"])
- # [bytes1, bytes1]
-
- # Iterate over keys
- for key in mongodb_store.yield_keys():
- print(key)
-
- # Delete keys
- mongodb_store.mdelete(["key1", "key2"])
- """
-
- def __init__(
- self,
- connection_string: str,
- db_name: str,
- collection_name: str,
- *,
- client_kwargs: Optional[dict] = None,
- ) -> None:
- """Initialize the MongoDBStore with a MongoDB connection string.
-
- Args:
- connection_string (str): MongoDB connection string
- db_name (str): name to use
- collection_name (str): collection name to use
- client_kwargs (dict): Keyword arguments to pass to the Mongo client
- """
- try:
- from pymongo import MongoClient
- except ImportError as e:
- raise ImportError(
- "The MongoDBStore requires the pymongo library to be "
- "installed. "
- "pip install pymongo"
- ) from e
-
- if not connection_string:
- raise ValueError("connection_string must be provided.")
- if not db_name:
- raise ValueError("db_name must be provided.")
- if not collection_name:
- raise ValueError("collection_name must be provided.")
-
- self.client: MongoClient = MongoClient(
- connection_string, **(client_kwargs or {})
- )
- self.collection = self.client[db_name][collection_name]
-
- def mget(self, keys: Sequence[str]) -> List[Optional[bytes]]:
- """Get the list of documents associated with the given keys.
-
- Args:
- keys (list[str]): A list of keys representing Document IDs..
-
- Returns:
- list[Document]: A list of Documents corresponding to the provided
- keys, where each Document is either retrieved successfully or
- represented as None if not found.
- """
- result = self.collection.find({"_id": {"$in": keys}})
- result_dict = {doc["_id"]: doc["value"] for doc in result}
- return [result_dict.get(key) for key in keys]
-
- def mset(self, key_value_pairs: Sequence[Tuple[str, bytes]]) -> None:
- """Set the given key-value pairs.
-
- Args:
- key_value_pairs (list[tuple[str, Document]]): A list of id-document
- pairs.
- """
- from pymongo import UpdateOne
-
- updates = [{"_id": k, "value": v} for k, v in key_value_pairs]
- self.collection.bulk_write(
- [UpdateOne({"_id": u["_id"]}, {"$set": u}, upsert=True) for u in updates]
- )
-
- def mdelete(self, keys: Sequence[str]) -> None:
- """Delete the given ids.
-
- Args:
- keys (list[str]): A list of keys representing Document IDs..
- """
- self.collection.delete_many({"_id": {"$in": keys}})
-
- def yield_keys(self, prefix: Optional[str] = None) -> Iterator[str]:
- """Yield keys in the store.
-
- Args:
- prefix (str): prefix of keys to retrieve.
- """
- if prefix is None:
- for doc in self.collection.find(projection=["_id"]):
- yield doc["_id"]
- else:
- for doc in self.collection.find(
- {"_id": {"$regex": f"^{prefix}"}}, projection=["_id"]
- ):
- yield doc["_id"]
-
-
-class MongoDBStore(BaseStore[str, Document]):
- """BaseStore implementation using MongoDB as the underlying store.
-
- Examples:
- Create a MongoDBStore instance and perform operations on it:
-
- .. code-block:: python
-
- # Instantiate the MongoDBStore with a MongoDB connection
- from langchain.storage import MongoDBStore
-
- mongo_conn_str = "mongodb://localhost:27017/"
- mongodb_store = MongoDBStore(mongo_conn_str, db_name="test-db",
- collection_name="test-collection")
-
- # Set values for keys
- doc1 = Document(...)
- doc2 = Document(...)
- mongodb_store.mset([("key1", doc1), ("key2", doc2)])
-
- # Get values for keys
- values = mongodb_store.mget(["key1", "key2"])
- # [doc1, doc2]
-
- # Iterate over keys
- for key in mongodb_store.yield_keys():
- print(key)
-
- # Delete keys
- mongodb_store.mdelete(["key1", "key2"])
- """
-
- def __init__(
- self,
- connection_string: str,
- db_name: str,
- collection_name: str,
- *,
- client_kwargs: Optional[dict] = None,
- ) -> None:
- """Initialize the MongoDBStore with a MongoDB connection string.
-
- Args:
- connection_string (str): MongoDB connection string
- db_name (str): name to use
- collection_name (str): collection name to use
- client_kwargs (dict): Keyword arguments to pass to the Mongo client
- """
- try:
- from pymongo import MongoClient
- except ImportError as e:
- raise ImportError(
- "The MongoDBStore requires the pymongo library to be "
- "installed. "
- "pip install pymongo"
- ) from e
-
- if not connection_string:
- raise ValueError("connection_string must be provided.")
- if not db_name:
- raise ValueError("db_name must be provided.")
- if not collection_name:
- raise ValueError("collection_name must be provided.")
-
- self.client: MongoClient = MongoClient(
- connection_string, **(client_kwargs or {})
- )
- self.collection = self.client[db_name][collection_name]
-
- def mget(self, keys: Sequence[str]) -> List[Optional[Document]]:
- """Get the list of documents associated with the given keys.
-
- Args:
- keys (list[str]): A list of keys representing Document IDs..
-
- Returns:
- list[Document]: A list of Documents corresponding to the provided
- keys, where each Document is either retrieved successfully or
- represented as None if not found.
- """
- result = self.collection.find({"_id": {"$in": keys}})
- result_dict = {doc["_id"]: Document(**doc["value"]) for doc in result}
- return [result_dict.get(key) for key in keys]
-
- def mset(self, key_value_pairs: Sequence[Tuple[str, Document]]) -> None:
- """Set the given key-value pairs.
-
- Args:
- key_value_pairs (list[tuple[str, Document]]): A list of id-document
- pairs.
- Returns:
- None
- """
- from pymongo import UpdateOne
-
- updates = [{"_id": k, "value": v.__dict__} for k, v in key_value_pairs]
- self.collection.bulk_write(
- [UpdateOne({"_id": u["_id"]}, {"$set": u}, upsert=True) for u in updates]
- )
-
- def mdelete(self, keys: Sequence[str]) -> None:
- """Delete the given ids.
-
- Args:
- keys (list[str]): A list of keys representing Document IDs..
- """
- self.collection.delete_many({"_id": {"$in": keys}})
-
- def yield_keys(self, prefix: Optional[str] = None) -> Iterator[str]:
- """Yield keys in the store.
-
- Args:
- prefix (str): prefix of keys to retrieve.
- """
- if prefix is None:
- for doc in self.collection.find(projection=["_id"]):
- yield doc["_id"]
- else:
- for doc in self.collection.find(
- {"_id": {"$regex": f"^{prefix}"}}, projection=["_id"]
- ):
- yield doc["_id"]
diff --git a/libs/community/langchain_community/storage/redis.py b/libs/community/langchain_community/storage/redis.py
deleted file mode 100644
index 2bf205d7d7..0000000000
--- a/libs/community/langchain_community/storage/redis.py
+++ /dev/null
@@ -1,144 +0,0 @@
-from typing import Any, Iterator, List, Optional, Sequence, Tuple, cast
-
-from langchain_core.stores import ByteStore
-
-from langchain_community.utilities.redis import get_client
-
-
-class RedisStore(ByteStore):
- """BaseStore implementation using Redis as the underlying store.
-
- Examples:
- Create a RedisStore instance and perform operations on it:
-
- .. code-block:: python
-
- # Instantiate the RedisStore with a Redis connection
- from langchain_community.storage import RedisStore
- from langchain_community.utilities.redis import get_client
-
- client = get_client('redis://localhost:6379')
- redis_store = RedisStore(client=client)
-
- # Set values for keys
- redis_store.mset([("key1", b"value1"), ("key2", b"value2")])
-
- # Get values for keys
- values = redis_store.mget(["key1", "key2"])
- # [b"value1", b"value2"]
-
- # Delete keys
- redis_store.mdelete(["key1"])
-
- # Iterate over keys
- for key in redis_store.yield_keys():
- print(key) # noqa: T201
- """
-
- def __init__(
- self,
- *,
- client: Any = None,
- redis_url: Optional[str] = None,
- client_kwargs: Optional[dict] = None,
- ttl: Optional[int] = None,
- namespace: Optional[str] = None,
- ) -> None:
- """Initialize the RedisStore with a Redis connection.
-
- Must provide either a Redis client or a redis_url with optional client_kwargs.
-
- Args:
- client: A Redis connection instance
- redis_url: redis url
- client_kwargs: Keyword arguments to pass to the Redis client
- ttl: time to expire keys in seconds if provided,
- if None keys will never expire
- namespace: if provided, all keys will be prefixed with this namespace
- """
- try:
- from redis import Redis
- except ImportError as e:
- raise ImportError(
- "The RedisStore requires the redis library to be installed. "
- "pip install redis"
- ) from e
-
- if client and (redis_url or client_kwargs):
- raise ValueError(
- "Either a Redis client or a redis_url with optional client_kwargs "
- "must be provided, but not both."
- )
-
- if not client and not redis_url:
- raise ValueError("Either a Redis client or a redis_url must be provided.")
-
- if client:
- if not isinstance(client, Redis):
- raise TypeError(
- f"Expected Redis client, got {type(client).__name__} instead."
- )
- _client = client
- else:
- if not redis_url:
- raise ValueError(
- "Either a Redis client or a redis_url must be provided."
- )
- _client = get_client(redis_url, **(client_kwargs or {}))
-
- self.client = _client
-
- if not isinstance(ttl, int) and ttl is not None:
- raise TypeError(f"Expected int or None, got {type(ttl)=} instead.")
-
- self.ttl = ttl
- self.namespace = namespace
-
- def _get_prefixed_key(self, key: str) -> str:
- """Get the key with the namespace prefix.
-
- Args:
- key (str): The original key.
-
- Returns:
- str: The key with the namespace prefix.
- """
- delimiter = "/"
- if self.namespace:
- return f"{self.namespace}{delimiter}{key}"
- return key
-
- def mget(self, keys: Sequence[str]) -> List[Optional[bytes]]:
- """Get the values associated with the given keys."""
- return cast(
- List[Optional[bytes]],
- self.client.mget([self._get_prefixed_key(key) for key in keys]),
- )
-
- def mset(self, key_value_pairs: Sequence[Tuple[str, bytes]]) -> None:
- """Set the given key-value pairs."""
- pipe = self.client.pipeline()
-
- for key, value in key_value_pairs:
- pipe.set(self._get_prefixed_key(key), value, ex=self.ttl)
- pipe.execute()
-
- def mdelete(self, keys: Sequence[str]) -> None:
- """Delete the given keys."""
- _keys = [self._get_prefixed_key(key) for key in keys]
- self.client.delete(*_keys)
-
- def yield_keys(self, *, prefix: Optional[str] = None) -> Iterator[str]:
- """Yield keys in the store."""
- if prefix:
- pattern = self._get_prefixed_key(prefix)
- else:
- pattern = self._get_prefixed_key("*")
- scan_iter = cast(Iterator[bytes], self.client.scan_iter(match=pattern))
- for key in scan_iter:
- decoded_key = key.decode("utf-8")
- if self.namespace:
- relative_key = decoded_key[len(self.namespace) + 1 :]
- yield relative_key
- else:
- yield decoded_key
diff --git a/libs/community/langchain_community/storage/sql.py b/libs/community/langchain_community/storage/sql.py
deleted file mode 100644
index c9652ae5f5..0000000000
--- a/libs/community/langchain_community/storage/sql.py
+++ /dev/null
@@ -1,295 +0,0 @@
-import contextlib
-from pathlib import Path
-from typing import (
- Any,
- AsyncGenerator,
- AsyncIterator,
- Dict,
- Generator,
- Iterator,
- List,
- Optional,
- Sequence,
- Tuple,
- Union,
- cast,
-)
-
-from langchain_core.stores import BaseStore
-from sqlalchemy import (
- LargeBinary,
- Text,
- and_,
- create_engine,
- delete,
- select,
-)
-from sqlalchemy.engine.base import Engine
-from sqlalchemy.ext.asyncio import (
- AsyncEngine,
- AsyncSession,
- create_async_engine,
-)
-from sqlalchemy.orm import (
- Mapped,
- Session,
- declarative_base,
- sessionmaker,
-)
-
-try:
- from sqlalchemy.ext.asyncio import async_sessionmaker
-except ImportError:
- # dummy for sqlalchemy < 2
- async_sessionmaker = type("async_sessionmaker", (type,), {}) # type: ignore[assignment,misc]
-
-Base = declarative_base()
-
-try:
- from sqlalchemy.orm import mapped_column
-
- class LangchainKeyValueStores(Base): # type: ignore[valid-type,misc]
- """Table used to save values."""
-
- # ATTENTION:
- # Prior to modifying this table, please determine whether
- # we should create migrations for this table to make sure
- # users do not experience data loss.
- __tablename__ = "langchain_key_value_stores"
-
- namespace: Mapped[str] = mapped_column(
- primary_key=True, index=True, nullable=False
- )
- key: Mapped[str] = mapped_column(primary_key=True, index=True, nullable=False)
- value = mapped_column(LargeBinary, index=False, nullable=False)
-
-except ImportError:
- # dummy for sqlalchemy < 2
- from sqlalchemy import Column
-
- class LangchainKeyValueStores(Base): # type: ignore[valid-type,misc,no-redef]
- """Table used to save values."""
-
- # ATTENTION:
- # Prior to modifying this table, please determine whether
- # we should create migrations for this table to make sure
- # users do not experience data loss.
- __tablename__ = "langchain_key_value_stores"
-
- namespace = Column(Text(), primary_key=True, index=True, nullable=False)
- key = Column(Text(), primary_key=True, index=True, nullable=False)
- value = Column(LargeBinary, index=False, nullable=False)
-
-
-def items_equal(x: Any, y: Any) -> bool:
- return x == y
-
-
-# This is a fix of original SQLStore.
-# This can will be removed when a PR will be merged.
-class SQLStore(BaseStore[str, bytes]):
- """BaseStore interface that works on an SQL database.
-
- Examples:
- Create a SQLStore instance and perform operations on it:
-
- .. code-block:: python
-
- from langchain_community.storage import SQLStore
-
- # Instantiate the SQLStore with the root path
- sql_store = SQLStore(namespace="test", db_url="sqlite://:memory:")
-
- # Set values for keys
- sql_store.mset([("key1", b"value1"), ("key2", b"value2")])
-
- # Get values for keys
- values = sql_store.mget(["key1", "key2"]) # Returns [b"value1", b"value2"]
-
- # Delete keys
- sql_store.mdelete(["key1"])
-
- # Iterate over keys
- for key in sql_store.yield_keys():
- print(key)
-
- """
-
- def __init__(
- self,
- *,
- namespace: str,
- db_url: Optional[Union[str, Path]] = None,
- engine: Optional[Union[Engine, AsyncEngine]] = None,
- engine_kwargs: Optional[Dict[str, Any]] = None,
- async_mode: Optional[bool] = None,
- ):
- if db_url is None and engine is None:
- raise ValueError("Must specify either db_url or engine")
-
- if db_url is not None and engine is not None:
- raise ValueError("Must specify either db_url or engine, not both")
-
- _engine: Union[Engine, AsyncEngine]
- if db_url:
- if async_mode is None:
- async_mode = False
- if async_mode:
- _engine = create_async_engine(
- url=str(db_url),
- **(engine_kwargs or {}),
- )
- else:
- _engine = create_engine(url=str(db_url), **(engine_kwargs or {}))
- elif engine:
- _engine = engine
-
- else:
- raise AssertionError("Something went wrong with configuration of engine.")
-
- _session_maker: Union[sessionmaker[Session], async_sessionmaker[AsyncSession]]
- if isinstance(_engine, AsyncEngine):
- self.async_mode = True
- _session_maker = async_sessionmaker(bind=_engine)
- else:
- self.async_mode = False
- _session_maker = sessionmaker(bind=_engine)
-
- self.engine = _engine
- self.dialect = _engine.dialect.name
- self.session_maker = _session_maker
- self.namespace = namespace
-
- def create_schema(self) -> None:
- Base.metadata.create_all(self.engine) # problem in sqlalchemy v1
- # sqlalchemy.exc.CompileError: (in table 'langchain_key_value_stores',
- # column 'namespace'): Can't generate DDL for NullType(); did you forget
- # to specify a type on this Column?
-
- async def acreate_schema(self) -> None:
- assert isinstance(self.engine, AsyncEngine)
- async with self.engine.begin() as session:
- await session.run_sync(Base.metadata.create_all)
-
- def drop(self) -> None:
- Base.metadata.drop_all(bind=self.engine.connect())
-
- async def amget(self, keys: Sequence[str]) -> List[Optional[bytes]]:
- assert isinstance(self.engine, AsyncEngine)
- result: Dict[str, bytes] = {}
- async with self._make_async_session() as session:
- stmt = select(LangchainKeyValueStores).filter(
- and_(
- LangchainKeyValueStores.key.in_(keys),
- LangchainKeyValueStores.namespace == self.namespace,
- )
- )
- for v in await session.scalars(stmt):
- result[v.key] = v.value
- return [result.get(key) for key in keys]
-
- def mget(self, keys: Sequence[str]) -> List[Optional[bytes]]:
- result = {}
-
- with self._make_sync_session() as session:
- stmt = select(LangchainKeyValueStores).filter(
- and_(
- LangchainKeyValueStores.key.in_(keys),
- LangchainKeyValueStores.namespace == self.namespace,
- )
- )
- for v in session.scalars(stmt):
- result[v.key] = v.value
- return [result.get(key) for key in keys]
-
- async def amset(self, key_value_pairs: Sequence[Tuple[str, bytes]]) -> None:
- async with self._make_async_session() as session:
- await self._amdelete([key for key, _ in key_value_pairs], session)
- session.add_all(
- [
- LangchainKeyValueStores(namespace=self.namespace, key=k, value=v)
- for k, v in key_value_pairs
- ]
- )
- await session.commit()
-
- def mset(self, key_value_pairs: Sequence[Tuple[str, bytes]]) -> None:
- values: Dict[str, bytes] = dict(key_value_pairs)
- with self._make_sync_session() as session:
- self._mdelete(list(values.keys()), session)
- session.add_all(
- [
- LangchainKeyValueStores(namespace=self.namespace, key=k, value=v)
- for k, v in values.items()
- ]
- )
- session.commit()
-
- def _mdelete(self, keys: Sequence[str], session: Session) -> None:
- stmt = delete(LangchainKeyValueStores).filter(
- and_(
- LangchainKeyValueStores.key.in_(keys),
- LangchainKeyValueStores.namespace == self.namespace,
- )
- )
- session.execute(stmt)
-
- async def _amdelete(self, keys: Sequence[str], session: AsyncSession) -> None:
- stmt = delete(LangchainKeyValueStores).filter(
- and_(
- LangchainKeyValueStores.key.in_(keys),
- LangchainKeyValueStores.namespace == self.namespace,
- )
- )
- await session.execute(stmt)
-
- def mdelete(self, keys: Sequence[str]) -> None:
- with self._make_sync_session() as session:
- self._mdelete(keys, session)
- session.commit()
-
- async def amdelete(self, keys: Sequence[str]) -> None:
- async with self._make_async_session() as session:
- await self._amdelete(keys, session)
- await session.commit()
-
- def yield_keys(self, *, prefix: Optional[str] = None) -> Iterator[str]:
- with self._make_sync_session() as session:
- for v in session.query(LangchainKeyValueStores).filter(
- LangchainKeyValueStores.namespace == self.namespace
- ):
- if str(v.key).startswith(prefix or ""):
- yield str(v.key)
- session.close()
-
- async def ayield_keys(self, *, prefix: Optional[str] = None) -> AsyncIterator[str]:
- async with self._make_async_session() as session:
- stmt = select(LangchainKeyValueStores).filter(
- LangchainKeyValueStores.namespace == self.namespace
- )
- for v in await session.scalars(stmt):
- if str(v.key).startswith(prefix or ""):
- yield str(v.key)
- await session.close()
-
- @contextlib.contextmanager
- def _make_sync_session(self) -> Generator[Session, None, None]:
- """Make an async session."""
- if self.async_mode:
- raise ValueError(
- "Attempting to use a sync method in when async mode is turned on. "
- "Please use the corresponding async method instead."
- )
- with cast(Session, self.session_maker()) as session:
- yield cast(Session, session)
-
- @contextlib.asynccontextmanager
- async def _make_async_session(self) -> AsyncGenerator[AsyncSession, None]:
- """Make an async session."""
- if not self.async_mode:
- raise ValueError(
- "Attempting to use an async method in when sync mode is turned on. "
- "Please use the corresponding async method instead."
- )
- async with cast(AsyncSession, self.session_maker()) as session:
- yield cast(AsyncSession, session)
diff --git a/libs/community/langchain_community/storage/upstash_redis.py b/libs/community/langchain_community/storage/upstash_redis.py
deleted file mode 100644
index ebe69c4dfb..0000000000
--- a/libs/community/langchain_community/storage/upstash_redis.py
+++ /dev/null
@@ -1,174 +0,0 @@
-from typing import Any, Iterator, List, Optional, Sequence, Tuple, cast
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.stores import BaseStore, ByteStore
-
-
-class _UpstashRedisStore(BaseStore[str, str]):
- """BaseStore implementation using Upstash Redis as the underlying store."""
-
- def __init__(
- self,
- *,
- client: Any = None,
- url: Optional[str] = None,
- token: Optional[str] = None,
- ttl: Optional[int] = None,
- namespace: Optional[str] = None,
- ) -> None:
- """Initialize the UpstashRedisStore with HTTP API.
-
- Must provide either an Upstash Redis client or a url.
-
- Args:
- client: An Upstash Redis instance
- url: UPSTASH_REDIS_REST_URL
- token: UPSTASH_REDIS_REST_TOKEN
- ttl: time to expire keys in seconds if provided,
- if None keys will never expire
- namespace: if provided, all keys will be prefixed with this namespace
- """
- try:
- from upstash_redis import Redis
- except ImportError as e:
- raise ImportError(
- "UpstashRedisStore requires the upstash_redis library to be installed. "
- "pip install upstash_redis"
- ) from e
-
- if client and url:
- raise ValueError(
- "Either an Upstash Redis client or a url must be provided, not both."
- )
-
- if client:
- if not isinstance(client, Redis):
- raise TypeError(
- f"Expected Upstash Redis client, got {type(client).__name__}."
- )
- _client = client
- else:
- if not url or not token:
- raise ValueError(
- "Either an Upstash Redis client or url and token must be provided."
- )
- _client = Redis(url=url, token=token)
-
- self.client = _client
-
- if not isinstance(ttl, int) and ttl is not None:
- raise TypeError(f"Expected int or None, got {type(ttl)} instead.")
-
- self.ttl = ttl
- self.namespace = namespace
-
- def _get_prefixed_key(self, key: str) -> str:
- """Get the key with the namespace prefix.
-
- Args:
- key (str): The original key.
-
- Returns:
- str: The key with the namespace prefix.
- """
- delimiter = "/"
- if self.namespace:
- return f"{self.namespace}{delimiter}{key}"
- return key
-
- def mget(self, keys: Sequence[str]) -> List[Optional[str]]:
- """Get the values associated with the given keys."""
-
- keys = [self._get_prefixed_key(key) for key in keys]
- return cast(
- List[Optional[str]],
- self.client.mget(*keys),
- )
-
- def mset(self, key_value_pairs: Sequence[Tuple[str, str]]) -> None:
- """Set the given key-value pairs."""
- for key, value in key_value_pairs:
- self.client.set(self._get_prefixed_key(key), value, ex=self.ttl)
-
- def mdelete(self, keys: Sequence[str]) -> None:
- """Delete the given keys."""
- _keys = [self._get_prefixed_key(key) for key in keys]
- self.client.delete(*_keys)
-
- def yield_keys(self, *, prefix: Optional[str] = None) -> Iterator[str]:
- """Yield keys in the store."""
- if prefix:
- pattern = self._get_prefixed_key(prefix)
- else:
- pattern = self._get_prefixed_key("*")
-
- cursor, keys = self.client.scan(0, match=pattern)
- for key in keys:
- if self.namespace:
- relative_key = key[len(self.namespace) + 1 :]
- yield relative_key
- else:
- yield key
-
- while cursor != 0:
- cursor, keys = self.client.scan(cursor, match=pattern)
- for key in keys:
- if self.namespace:
- relative_key = key[len(self.namespace) + 1 :]
- yield relative_key
- else:
- yield key
-
-
-@deprecated("0.0.1", alternative="UpstashRedisByteStore")
-class UpstashRedisStore(_UpstashRedisStore):
- """
- BaseStore implementation using Upstash Redis
- as the underlying store to store strings.
-
- Deprecated in favor of the more generic UpstashRedisByteStore.
- """
-
-
-class UpstashRedisByteStore(ByteStore):
- """
- BaseStore implementation using Upstash Redis
- as the underlying store to store raw bytes.
- """
-
- def __init__(
- self,
- *,
- client: Any = None,
- url: Optional[str] = None,
- token: Optional[str] = None,
- ttl: Optional[int] = None,
- namespace: Optional[str] = None,
- ) -> None:
- self.underlying_store = _UpstashRedisStore(
- client=client, url=url, token=token, ttl=ttl, namespace=namespace
- )
-
- def mget(self, keys: Sequence[str]) -> List[Optional[bytes]]:
- """Get the values associated with the given keys."""
- return [
- value.encode("utf-8") if value is not None else None
- for value in self.underlying_store.mget(keys)
- ]
-
- def mset(self, key_value_pairs: Sequence[Tuple[str, bytes]]) -> None:
- """Set the given key-value pairs."""
- self.underlying_store.mset(
- [
- (k, v.decode("utf-8")) if v is not None else None
- for k, v in key_value_pairs
- ]
- )
-
- def mdelete(self, keys: Sequence[str]) -> None:
- """Delete the given keys."""
- self.underlying_store.mdelete(keys)
-
- def yield_keys(self, *, prefix: Optional[str] = None) -> Iterator[str]:
- """Yield keys in the store."""
- yield from self.underlying_store.yield_keys(prefix=prefix)
diff --git a/libs/community/langchain_community/tools/__init__.py b/libs/community/langchain_community/tools/__init__.py
deleted file mode 100644
index de486cfbb3..0000000000
--- a/libs/community/langchain_community/tools/__init__.py
+++ /dev/null
@@ -1,664 +0,0 @@
-"""**Tools** are classes that an Agent uses to interact with the world.
-
-Each tool has a **description**. Agent uses the description to choose the right
-tool for the job.
-
-**Class hierarchy:**
-
-.. code-block::
-
- ToolMetaclass --> BaseTool --> Tool # Examples: AIPluginTool, BaseGraphQLTool
- # Examples: BraveSearch, HumanInputRun
-
-**Main helpers:**
-
-.. code-block::
-
- CallbackManagerForToolRun, AsyncCallbackManagerForToolRun
-"""
-
-import importlib
-from typing import TYPE_CHECKING, Any
-
-if TYPE_CHECKING:
- from langchain_core.tools import (
- BaseTool as BaseTool,
- )
- from langchain_core.tools import (
- StructuredTool as StructuredTool,
- )
- from langchain_core.tools import (
- Tool as Tool,
- )
- from langchain_core.tools.convert import tool as tool
-
- from langchain_community.tools.ainetwork.app import (
- AINAppOps,
- )
- from langchain_community.tools.ainetwork.owner import (
- AINOwnerOps,
- )
- from langchain_community.tools.ainetwork.rule import (
- AINRuleOps,
- )
- from langchain_community.tools.ainetwork.transfer import (
- AINTransfer,
- )
- from langchain_community.tools.ainetwork.value import (
- AINValueOps,
- )
- from langchain_community.tools.arxiv.tool import (
- ArxivQueryRun,
- )
- from langchain_community.tools.asknews.tool import (
- AskNewsSearch,
- )
- from langchain_community.tools.azure_ai_services import (
- AzureAiServicesDocumentIntelligenceTool,
- AzureAiServicesImageAnalysisTool,
- AzureAiServicesSpeechToTextTool,
- AzureAiServicesTextAnalyticsForHealthTool,
- AzureAiServicesTextToSpeechTool,
- )
- from langchain_community.tools.azure_cognitive_services import (
- AzureCogsFormRecognizerTool,
- AzureCogsImageAnalysisTool,
- AzureCogsSpeech2TextTool,
- AzureCogsText2SpeechTool,
- AzureCogsTextAnalyticsHealthTool,
- )
- from langchain_community.tools.bearly.tool import (
- BearlyInterpreterTool,
- )
- from langchain_community.tools.bing_search.tool import (
- BingSearchResults,
- BingSearchRun,
- )
- from langchain_community.tools.brave_search.tool import (
- BraveSearch,
- )
- from langchain_community.tools.cassandra_database.tool import (
- GetSchemaCassandraDatabaseTool, # noqa: F401
- GetTableDataCassandraDatabaseTool, # noqa: F401
- QueryCassandraDatabaseTool, # noqa: F401
- )
- from langchain_community.tools.cogniswitch.tool import (
- CogniswitchKnowledgeRequest,
- CogniswitchKnowledgeSourceFile,
- CogniswitchKnowledgeSourceURL,
- CogniswitchKnowledgeStatus,
- )
- from langchain_community.tools.connery import (
- ConneryAction,
- )
- from langchain_community.tools.convert_to_openai import (
- format_tool_to_openai_function,
- )
- from langchain_community.tools.dataherald import DataheraldTextToSQL
- from langchain_community.tools.ddg_search.tool import (
- DuckDuckGoSearchResults,
- DuckDuckGoSearchRun,
- )
- from langchain_community.tools.e2b_data_analysis.tool import (
- E2BDataAnalysisTool,
- )
- from langchain_community.tools.edenai import (
- EdenAiExplicitImageTool,
- EdenAiObjectDetectionTool,
- EdenAiParsingIDTool,
- EdenAiParsingInvoiceTool,
- EdenAiSpeechToTextTool,
- EdenAiTextModerationTool,
- EdenAiTextToSpeechTool,
- EdenaiTool,
- )
- from langchain_community.tools.eleven_labs.text2speech import (
- ElevenLabsText2SpeechTool,
- )
- from langchain_community.tools.file_management import (
- CopyFileTool,
- DeleteFileTool,
- FileSearchTool,
- ListDirectoryTool,
- MoveFileTool,
- ReadFileTool,
- WriteFileTool,
- )
- from langchain_community.tools.financial_datasets.balance_sheets import (
- BalanceSheets,
- )
- from langchain_community.tools.financial_datasets.cash_flow_statements import (
- CashFlowStatements,
- )
- from langchain_community.tools.financial_datasets.income_statements import (
- IncomeStatements,
- )
- from langchain_community.tools.gmail import (
- GmailCreateDraft,
- GmailGetMessage,
- GmailGetThread,
- GmailSearch,
- GmailSendMessage,
- )
- from langchain_community.tools.google_books import (
- GoogleBooksQueryRun,
- )
- from langchain_community.tools.google_cloud.texttospeech import (
- GoogleCloudTextToSpeechTool,
- )
- from langchain_community.tools.google_places.tool import (
- GooglePlacesTool,
- )
- from langchain_community.tools.google_search.tool import (
- GoogleSearchResults,
- GoogleSearchRun,
- )
- from langchain_community.tools.google_serper.tool import (
- GoogleSerperResults,
- GoogleSerperRun,
- )
- from langchain_community.tools.graphql.tool import (
- BaseGraphQLTool,
- )
- from langchain_community.tools.human.tool import (
- HumanInputRun,
- )
- from langchain_community.tools.ifttt import (
- IFTTTWebhook,
- )
- from langchain_community.tools.interaction.tool import (
- StdInInquireTool,
- )
- from langchain_community.tools.jina_search.tool import JinaSearch
- from langchain_community.tools.jira.tool import (
- JiraAction,
- )
- from langchain_community.tools.json.tool import (
- JsonGetValueTool,
- JsonListKeysTool,
- )
- from langchain_community.tools.merriam_webster.tool import (
- MerriamWebsterQueryRun,
- )
- from langchain_community.tools.metaphor_search import (
- MetaphorSearchResults,
- )
- from langchain_community.tools.mojeek_search.tool import (
- MojeekSearch,
- )
- from langchain_community.tools.nasa.tool import (
- NasaAction,
- )
- from langchain_community.tools.office365.create_draft_message import (
- O365CreateDraftMessage,
- )
- from langchain_community.tools.office365.events_search import (
- O365SearchEvents,
- )
- from langchain_community.tools.office365.messages_search import (
- O365SearchEmails,
- )
- from langchain_community.tools.office365.send_event import (
- O365SendEvent,
- )
- from langchain_community.tools.office365.send_message import (
- O365SendMessage,
- )
- from langchain_community.tools.office365.utils import (
- authenticate,
- )
- from langchain_community.tools.openapi.utils.api_models import (
- APIOperation,
- )
- from langchain_community.tools.openapi.utils.openapi_utils import (
- OpenAPISpec,
- )
- from langchain_community.tools.openweathermap.tool import (
- OpenWeatherMapQueryRun,
- )
- from langchain_community.tools.playwright import (
- ClickTool,
- CurrentWebPageTool,
- ExtractHyperlinksTool,
- ExtractTextTool,
- GetElementsTool,
- NavigateBackTool,
- NavigateTool,
- )
- from langchain_community.tools.plugin import (
- AIPluginTool,
- )
- from langchain_community.tools.polygon.aggregates import (
- PolygonAggregates,
- )
- from langchain_community.tools.polygon.financials import (
- PolygonFinancials,
- )
- from langchain_community.tools.polygon.last_quote import (
- PolygonLastQuote,
- )
- from langchain_community.tools.polygon.ticker_news import (
- PolygonTickerNews,
- )
- from langchain_community.tools.powerbi.tool import (
- InfoPowerBITool,
- ListPowerBITool,
- QueryPowerBITool,
- )
- from langchain_community.tools.pubmed.tool import (
- PubmedQueryRun,
- )
- from langchain_community.tools.reddit_search.tool import (
- RedditSearchRun,
- RedditSearchSchema,
- )
- from langchain_community.tools.requests.tool import (
- BaseRequestsTool,
- RequestsDeleteTool,
- RequestsGetTool,
- RequestsPatchTool,
- RequestsPostTool,
- RequestsPutTool,
- )
- from langchain_community.tools.scenexplain.tool import (
- SceneXplainTool,
- )
- from langchain_community.tools.searchapi.tool import (
- SearchAPIResults,
- SearchAPIRun,
- )
- from langchain_community.tools.searx_search.tool import (
- SearxSearchResults,
- SearxSearchRun,
- )
- from langchain_community.tools.shell.tool import (
- ShellTool,
- )
- from langchain_community.tools.slack.get_channel import (
- SlackGetChannel,
- )
- from langchain_community.tools.slack.get_message import (
- SlackGetMessage,
- )
- from langchain_community.tools.slack.schedule_message import (
- SlackScheduleMessage,
- )
- from langchain_community.tools.slack.send_message import (
- SlackSendMessage,
- )
- from langchain_community.tools.sleep.tool import (
- SleepTool,
- )
- from langchain_community.tools.spark_sql.tool import (
- BaseSparkSQLTool,
- InfoSparkSQLTool,
- ListSparkSQLTool,
- QueryCheckerTool,
- QuerySparkSQLTool,
- )
- from langchain_community.tools.sql_database.tool import (
- BaseSQLDatabaseTool,
- InfoSQLDatabaseTool,
- ListSQLDatabaseTool,
- QuerySQLCheckerTool,
- QuerySQLDataBaseTool,
- QuerySQLDatabaseTool,
- )
- from langchain_community.tools.stackexchange.tool import (
- StackExchangeTool,
- )
- from langchain_community.tools.steam.tool import (
- SteamWebAPIQueryRun,
- )
- from langchain_community.tools.steamship_image_generation import (
- SteamshipImageGenerationTool,
- )
- from langchain_community.tools.tavily_search import (
- TavilyAnswer,
- TavilySearchResults,
- )
- from langchain_community.tools.vectorstore.tool import (
- VectorStoreQATool,
- VectorStoreQAWithSourcesTool,
- )
- from langchain_community.tools.wikipedia.tool import (
- WikipediaQueryRun,
- )
- from langchain_community.tools.wolfram_alpha.tool import (
- WolframAlphaQueryRun,
- )
- from langchain_community.tools.yahoo_finance_news import (
- YahooFinanceNewsTool,
- )
- from langchain_community.tools.you.tool import (
- YouSearchTool,
- )
- from langchain_community.tools.youtube.search import (
- YouTubeSearchTool,
- )
- from langchain_community.tools.zapier.tool import (
- ZapierNLAListActions,
- ZapierNLARunAction,
- )
- from langchain_community.tools.zenguard.tool import (
- Detector,
- ZenGuardInput,
- ZenGuardTool,
- )
-
-__all__ = [
- "BaseTool",
- "Tool",
- "tool",
- "StructuredTool",
- "AINAppOps",
- "AINOwnerOps",
- "AINRuleOps",
- "AINTransfer",
- "AINValueOps",
- "AIPluginTool",
- "APIOperation",
- "ArxivQueryRun",
- "AskNewsSearch",
- "AzureAiServicesDocumentIntelligenceTool",
- "AzureAiServicesImageAnalysisTool",
- "AzureAiServicesSpeechToTextTool",
- "AzureAiServicesTextAnalyticsForHealthTool",
- "AzureAiServicesTextToSpeechTool",
- "AzureCogsFormRecognizerTool",
- "AzureCogsImageAnalysisTool",
- "AzureCogsSpeech2TextTool",
- "AzureCogsText2SpeechTool",
- "AzureCogsTextAnalyticsHealthTool",
- "BalanceSheets",
- "BaseGraphQLTool",
- "BaseRequestsTool",
- "BaseSQLDatabaseTool",
- "BaseSparkSQLTool",
- "BearlyInterpreterTool",
- "BingSearchResults",
- "BingSearchRun",
- "BraveSearch",
- "CashFlowStatements",
- "ClickTool",
- "CogniswitchKnowledgeRequest",
- "CogniswitchKnowledgeSourceFile",
- "CogniswitchKnowledgeSourceURL",
- "CogniswitchKnowledgeStatus",
- "ConneryAction",
- "CopyFileTool",
- "CurrentWebPageTool",
- "DeleteFileTool",
- "DataheraldTextToSQL",
- "DuckDuckGoSearchResults",
- "DuckDuckGoSearchRun",
- "E2BDataAnalysisTool",
- "EdenAiExplicitImageTool",
- "EdenAiObjectDetectionTool",
- "EdenAiParsingIDTool",
- "EdenAiParsingInvoiceTool",
- "EdenAiSpeechToTextTool",
- "EdenAiTextModerationTool",
- "EdenAiTextToSpeechTool",
- "EdenaiTool",
- "ElevenLabsText2SpeechTool",
- "ExtractHyperlinksTool",
- "ExtractTextTool",
- "FileSearchTool",
- "GetElementsTool",
- "GmailCreateDraft",
- "GmailGetMessage",
- "GmailGetThread",
- "GmailSearch",
- "GmailSendMessage",
- "GoogleBooksQueryRun",
- "GoogleCloudTextToSpeechTool",
- "GooglePlacesTool",
- "GoogleSearchResults",
- "GoogleSearchRun",
- "GoogleSerperResults",
- "GoogleSerperRun",
- "HumanInputRun",
- "IFTTTWebhook",
- "IncomeStatements",
- "InfoPowerBITool",
- "InfoSQLDatabaseTool",
- "InfoSparkSQLTool",
- "JiraAction",
- "JinaSearch",
- "JsonGetValueTool",
- "JsonListKeysTool",
- "ListDirectoryTool",
- "ListPowerBITool",
- "ListSQLDatabaseTool",
- "ListSparkSQLTool",
- "MerriamWebsterQueryRun",
- "MetaphorSearchResults",
- "MojeekSearch",
- "MoveFileTool",
- "NasaAction",
- "NavigateBackTool",
- "NavigateTool",
- "O365CreateDraftMessage",
- "O365SearchEmails",
- "O365SearchEvents",
- "O365SendEvent",
- "O365SendMessage",
- "OpenAPISpec",
- "OpenWeatherMapQueryRun",
- "PolygonAggregates",
- "PolygonFinancials",
- "PolygonLastQuote",
- "PolygonTickerNews",
- "PubmedQueryRun",
- "QueryCheckerTool",
- "QueryPowerBITool",
- "QuerySQLCheckerTool",
- "QuerySQLDatabaseTool",
- "QuerySQLDataBaseTool", # Legacy, kept for backwards compatibility.
- "QuerySparkSQLTool",
- "ReadFileTool",
- "RedditSearchRun",
- "RedditSearchSchema",
- "RequestsDeleteTool",
- "RequestsGetTool",
- "RequestsPatchTool",
- "RequestsPostTool",
- "RequestsPutTool",
- "SceneXplainTool",
- "SearchAPIResults",
- "SearchAPIRun",
- "SearxSearchResults",
- "SearxSearchRun",
- "ShellTool",
- "SlackGetChannel",
- "SlackGetMessage",
- "SlackScheduleMessage",
- "SlackSendMessage",
- "SleepTool",
- "StackExchangeTool",
- "StdInInquireTool",
- "SteamWebAPIQueryRun",
- "SteamshipImageGenerationTool",
- "TavilyAnswer",
- "TavilySearchResults",
- "VectorStoreQATool",
- "VectorStoreQAWithSourcesTool",
- "WikipediaQueryRun",
- "WolframAlphaQueryRun",
- "WriteFileTool",
- "YahooFinanceNewsTool",
- "YouSearchTool",
- "YouTubeSearchTool",
- "ZapierNLAListActions",
- "ZapierNLARunAction",
- "Detector",
- "ZenGuardInput",
- "ZenGuardTool",
- "authenticate",
- "format_tool_to_openai_function",
-]
-
-# Used for internal purposes
-_DEPRECATED_TOOLS = {"PythonAstREPLTool", "PythonREPLTool"}
-
-_module_lookup = {
- "AINAppOps": "langchain_community.tools.ainetwork.app",
- "AINOwnerOps": "langchain_community.tools.ainetwork.owner",
- "AINRuleOps": "langchain_community.tools.ainetwork.rule",
- "AINTransfer": "langchain_community.tools.ainetwork.transfer",
- "AINValueOps": "langchain_community.tools.ainetwork.value",
- "AIPluginTool": "langchain_community.tools.plugin",
- "APIOperation": "langchain_community.tools.openapi.utils.api_models",
- "ArxivQueryRun": "langchain_community.tools.arxiv.tool",
- "AskNewsSearch": "langchain_community.tools.asknews.tool",
- "AzureAiServicesDocumentIntelligenceTool": "langchain_community.tools.azure_ai_services", # noqa: E501
- "AzureAiServicesImageAnalysisTool": "langchain_community.tools.azure_ai_services",
- "AzureAiServicesSpeechToTextTool": "langchain_community.tools.azure_ai_services",
- "AzureAiServicesTextToSpeechTool": "langchain_community.tools.azure_ai_services",
- "AzureAiServicesTextAnalyticsForHealthTool": "langchain_community.tools.azure_ai_services", # noqa: E501
- "AzureCogsFormRecognizerTool": "langchain_community.tools.azure_cognitive_services",
- "AzureCogsImageAnalysisTool": "langchain_community.tools.azure_cognitive_services",
- "AzureCogsSpeech2TextTool": "langchain_community.tools.azure_cognitive_services",
- "AzureCogsText2SpeechTool": "langchain_community.tools.azure_cognitive_services",
- "AzureCogsTextAnalyticsHealthTool": "langchain_community.tools.azure_cognitive_services", # noqa: E501
- "BalanceSheets": "langchain_community.tools.financial_datasets.balance_sheets",
- "BaseGraphQLTool": "langchain_community.tools.graphql.tool",
- "BaseRequestsTool": "langchain_community.tools.requests.tool",
- "BaseSQLDatabaseTool": "langchain_community.tools.sql_database.tool",
- "BaseSparkSQLTool": "langchain_community.tools.spark_sql.tool",
- "BaseTool": "langchain_core.tools",
- "BearlyInterpreterTool": "langchain_community.tools.bearly.tool",
- "BingSearchResults": "langchain_community.tools.bing_search.tool",
- "BingSearchRun": "langchain_community.tools.bing_search.tool",
- "BraveSearch": "langchain_community.tools.brave_search.tool",
- "CashFlowStatements": "langchain_community.tools.financial_datasets.cash_flow_statements", # noqa: E501
- "ClickTool": "langchain_community.tools.playwright",
- "CogniswitchKnowledgeRequest": "langchain_community.tools.cogniswitch.tool",
- "CogniswitchKnowledgeSourceFile": "langchain_community.tools.cogniswitch.tool",
- "CogniswitchKnowledgeSourceURL": "langchain_community.tools.cogniswitch.tool",
- "CogniswitchKnowledgeStatus": "langchain_community.tools.cogniswitch.tool",
- "ConneryAction": "langchain_community.tools.connery",
- "CopyFileTool": "langchain_community.tools.file_management",
- "CurrentWebPageTool": "langchain_community.tools.playwright",
- "DataheraldTextToSQL": "langchain_community.tools.dataherald.tool",
- "DeleteFileTool": "langchain_community.tools.file_management",
- "Detector": "langchain_community.tools.zenguard.tool",
- "DuckDuckGoSearchResults": "langchain_community.tools.ddg_search.tool",
- "DuckDuckGoSearchRun": "langchain_community.tools.ddg_search.tool",
- "E2BDataAnalysisTool": "langchain_community.tools.e2b_data_analysis.tool",
- "EdenAiExplicitImageTool": "langchain_community.tools.edenai",
- "EdenAiObjectDetectionTool": "langchain_community.tools.edenai",
- "EdenAiParsingIDTool": "langchain_community.tools.edenai",
- "EdenAiParsingInvoiceTool": "langchain_community.tools.edenai",
- "EdenAiSpeechToTextTool": "langchain_community.tools.edenai",
- "EdenAiTextModerationTool": "langchain_community.tools.edenai",
- "EdenAiTextToSpeechTool": "langchain_community.tools.edenai",
- "EdenaiTool": "langchain_community.tools.edenai",
- "ElevenLabsText2SpeechTool": "langchain_community.tools.eleven_labs.text2speech",
- "ExtractHyperlinksTool": "langchain_community.tools.playwright",
- "ExtractTextTool": "langchain_community.tools.playwright",
- "FileSearchTool": "langchain_community.tools.file_management",
- "GetElementsTool": "langchain_community.tools.playwright",
- "GmailCreateDraft": "langchain_community.tools.gmail",
- "GmailGetMessage": "langchain_community.tools.gmail",
- "GmailGetThread": "langchain_community.tools.gmail",
- "GmailSearch": "langchain_community.tools.gmail",
- "GmailSendMessage": "langchain_community.tools.gmail",
- "GoogleBooksQueryRun": "langchain_community.tools.google_books",
- "GoogleCloudTextToSpeechTool": "langchain_community.tools.google_cloud.texttospeech", # noqa: E501
- "GooglePlacesTool": "langchain_community.tools.google_places.tool",
- "GoogleSearchResults": "langchain_community.tools.google_search.tool",
- "GoogleSearchRun": "langchain_community.tools.google_search.tool",
- "GoogleSerperResults": "langchain_community.tools.google_serper.tool",
- "GoogleSerperRun": "langchain_community.tools.google_serper.tool",
- "HumanInputRun": "langchain_community.tools.human.tool",
- "IFTTTWebhook": "langchain_community.tools.ifttt",
- "IncomeStatements": "langchain_community.tools.financial_datasets.income_statements", # noqa: E501
- "InfoPowerBITool": "langchain_community.tools.powerbi.tool",
- "InfoSQLDatabaseTool": "langchain_community.tools.sql_database.tool",
- "InfoSparkSQLTool": "langchain_community.tools.spark_sql.tool",
- "JiraAction": "langchain_community.tools.jira.tool",
- "JinaSearch": "langchain_community.tools.jina_search.tool",
- "JsonGetValueTool": "langchain_community.tools.json.tool",
- "JsonListKeysTool": "langchain_community.tools.json.tool",
- "ListDirectoryTool": "langchain_community.tools.file_management",
- "ListPowerBITool": "langchain_community.tools.powerbi.tool",
- "ListSQLDatabaseTool": "langchain_community.tools.sql_database.tool",
- "ListSparkSQLTool": "langchain_community.tools.spark_sql.tool",
- "MerriamWebsterQueryRun": "langchain_community.tools.merriam_webster.tool",
- "MetaphorSearchResults": "langchain_community.tools.metaphor_search",
- "MojeekSearch": "langchain_community.tools.mojeek_search.tool",
- "MoveFileTool": "langchain_community.tools.file_management",
- "NasaAction": "langchain_community.tools.nasa.tool",
- "NavigateBackTool": "langchain_community.tools.playwright",
- "NavigateTool": "langchain_community.tools.playwright",
- "O365CreateDraftMessage": "langchain_community.tools.office365.create_draft_message", # noqa: E501
- "O365SearchEmails": "langchain_community.tools.office365.messages_search",
- "O365SearchEvents": "langchain_community.tools.office365.events_search",
- "O365SendEvent": "langchain_community.tools.office365.send_event",
- "O365SendMessage": "langchain_community.tools.office365.send_message",
- "OpenAPISpec": "langchain_community.tools.openapi.utils.openapi_utils",
- "OpenWeatherMapQueryRun": "langchain_community.tools.openweathermap.tool",
- "PolygonAggregates": "langchain_community.tools.polygon.aggregates",
- "PolygonFinancials": "langchain_community.tools.polygon.financials",
- "PolygonLastQuote": "langchain_community.tools.polygon.last_quote",
- "PolygonTickerNews": "langchain_community.tools.polygon.ticker_news",
- "PubmedQueryRun": "langchain_community.tools.pubmed.tool",
- "QueryCheckerTool": "langchain_community.tools.spark_sql.tool",
- "QueryPowerBITool": "langchain_community.tools.powerbi.tool",
- "QuerySQLCheckerTool": "langchain_community.tools.sql_database.tool",
- "QuerySQLDatabaseTool": "langchain_community.tools.sql_database.tool",
- # Legacy, kept for backwards compatibility.
- "QuerySQLDataBaseTool": "langchain_community.tools.sql_database.tool",
- "QuerySparkSQLTool": "langchain_community.tools.spark_sql.tool",
- "ReadFileTool": "langchain_community.tools.file_management",
- "RedditSearchRun": "langchain_community.tools.reddit_search.tool",
- "RedditSearchSchema": "langchain_community.tools.reddit_search.tool",
- "RequestsDeleteTool": "langchain_community.tools.requests.tool",
- "RequestsGetTool": "langchain_community.tools.requests.tool",
- "RequestsPatchTool": "langchain_community.tools.requests.tool",
- "RequestsPostTool": "langchain_community.tools.requests.tool",
- "RequestsPutTool": "langchain_community.tools.requests.tool",
- "SceneXplainTool": "langchain_community.tools.scenexplain.tool",
- "SearchAPIResults": "langchain_community.tools.searchapi.tool",
- "SearchAPIRun": "langchain_community.tools.searchapi.tool",
- "SearxSearchResults": "langchain_community.tools.searx_search.tool",
- "SearxSearchRun": "langchain_community.tools.searx_search.tool",
- "ShellTool": "langchain_community.tools.shell.tool",
- "SlackGetChannel": "langchain_community.tools.slack.get_channel",
- "SlackGetMessage": "langchain_community.tools.slack.get_message",
- "SlackScheduleMessage": "langchain_community.tools.slack.schedule_message",
- "SlackSendMessage": "langchain_community.tools.slack.send_message",
- "SleepTool": "langchain_community.tools.sleep.tool",
- "StackExchangeTool": "langchain_community.tools.stackexchange.tool",
- "StdInInquireTool": "langchain_community.tools.interaction.tool",
- "SteamWebAPIQueryRun": "langchain_community.tools.steam.tool",
- "SteamshipImageGenerationTool": "langchain_community.tools.steamship_image_generation", # noqa: E501
- "StructuredTool": "langchain_core.tools",
- "TavilyAnswer": "langchain_community.tools.tavily_search",
- "TavilySearchResults": "langchain_community.tools.tavily_search",
- "Tool": "langchain_core.tools",
- "VectorStoreQATool": "langchain_community.tools.vectorstore.tool",
- "VectorStoreQAWithSourcesTool": "langchain_community.tools.vectorstore.tool",
- "WikipediaQueryRun": "langchain_community.tools.wikipedia.tool",
- "WolframAlphaQueryRun": "langchain_community.tools.wolfram_alpha.tool",
- "WriteFileTool": "langchain_community.tools.file_management",
- "YahooFinanceNewsTool": "langchain_community.tools.yahoo_finance_news",
- "YouSearchTool": "langchain_community.tools.you.tool",
- "YouTubeSearchTool": "langchain_community.tools.youtube.search",
- "ZapierNLAListActions": "langchain_community.tools.zapier.tool",
- "ZapierNLARunAction": "langchain_community.tools.zapier.tool",
- "ZenGuardInput": "langchain_community.tools.zenguard.tool",
- "ZenGuardTool": "langchain_community.tools.zenguard.tool",
- "authenticate": "langchain_community.tools.office365.utils",
- "format_tool_to_openai_function": "langchain_community.tools.convert_to_openai",
- "tool": "langchain_core.tools",
-}
-
-
-def __getattr__(name: str) -> Any:
- if name in _module_lookup:
- module = importlib.import_module(_module_lookup[name])
- return getattr(module, name)
- raise AttributeError(f"module {__name__} has no attribute {name}")
diff --git a/libs/community/langchain_community/tools/ainetwork/__init__.py b/libs/community/langchain_community/tools/ainetwork/__init__.py
deleted file mode 100644
index e69de29bb2..0000000000
diff --git a/libs/community/langchain_community/tools/ainetwork/app.py b/libs/community/langchain_community/tools/ainetwork/app.py
deleted file mode 100644
index 8175a210b7..0000000000
--- a/libs/community/langchain_community/tools/ainetwork/app.py
+++ /dev/null
@@ -1,102 +0,0 @@
-import builtins
-import json
-from enum import Enum
-from typing import List, Optional, Type, Union
-
-from langchain_core.callbacks import AsyncCallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.ainetwork.base import AINBaseTool
-
-
-class AppOperationType(str, Enum):
- """Type of app operation as enumerator."""
-
- SET_ADMIN = "SET_ADMIN"
- GET_CONFIG = "GET_CONFIG"
-
-
-class AppSchema(BaseModel):
- """Schema for app operations."""
-
- type: AppOperationType = Field(...)
- appName: str = Field(..., description="Name of the application on the blockchain")
- address: Optional[Union[str, List[str]]] = Field(
- None,
- description=(
- "A single address or a list of addresses. Default: current session's "
- "address"
- ),
- )
-
-
-class AINAppOps(AINBaseTool):
- """Tool for app operations."""
-
- name: str = "AINappOps"
- description: str = """
-Create an app in the AINetwork Blockchain database by creating the /apps/ path.
-An address set as `admin` can grant `owner` rights to other addresses (refer to `AINownerOps` for more details).
-Also, `admin` is initialized to have all `owner` permissions and `rule` allowed for that path.
-
-## appName Rule
-- [a-z_0-9]+
-
-## address Rules
-- 0x[0-9a-fA-F]{40}
-- Defaults to the current session's address
-- Multiple addresses can be specified if needed
-
-## SET_ADMIN Example 1
-- type: SET_ADMIN
-- appName: ain_project
-
-### Result:
-1. Path /apps/ain_project created.
-2. Current session's address registered as admin.
-
-## SET_ADMIN Example 2
-- type: SET_ADMIN
-- appName: test_project
-- address: [, ]
-
-### Result:
-1. Path /apps/test_project created.
-2. and registered as admin.
-
-""" # noqa: E501
- args_schema: Type[BaseModel] = AppSchema
-
- async def _arun(
- self,
- type: AppOperationType,
- appName: str,
- address: Optional[Union[str, List[str]]] = None,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- from ain.types import ValueOnlyTransactionInput
- from ain.utils import getTimestamp
-
- try:
- if type is AppOperationType.SET_ADMIN:
- if address is None:
- address = self.interface.wallet.defaultAccount.address
- if isinstance(address, str):
- address = [address]
-
- res = await self.interface.db.ref(
- f"/manage_app/{appName}/create/{getTimestamp()}"
- ).setValue(
- transactionInput=ValueOnlyTransactionInput(
- value={"admin": {address: True for address in address}}
- )
- )
- elif type is AppOperationType.GET_CONFIG:
- res = await self.interface.db.ref(
- f"/manage_app/{appName}/config"
- ).getValue()
- else:
- raise ValueError(f"Unsupported 'type': {type}.")
- return json.dumps(res, ensure_ascii=False)
- except Exception as e:
- return f"{builtins.type(e).__name__}: {str(e)}"
diff --git a/libs/community/langchain_community/tools/ainetwork/base.py b/libs/community/langchain_community/tools/ainetwork/base.py
deleted file mode 100644
index 00e4fc7f7a..0000000000
--- a/libs/community/langchain_community/tools/ainetwork/base.py
+++ /dev/null
@@ -1,73 +0,0 @@
-from __future__ import annotations
-
-import asyncio
-import threading
-from enum import Enum
-from typing import TYPE_CHECKING, Any, Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import Field
-
-from langchain_community.tools.ainetwork.utils import authenticate
-
-if TYPE_CHECKING:
- from ain.ain import Ain
-
-
-class OperationType(str, Enum):
- """Type of operation as enumerator."""
-
- SET = "SET"
- GET = "GET"
-
-
-class AINBaseTool(BaseTool):
- """Base class for the AINetwork tools."""
-
- interface: Ain = Field(default_factory=authenticate)
- """The interface object for the AINetwork Blockchain."""
-
- def _run(
- self,
- *args: Any,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- **kwargs: Any,
- ) -> str:
- try:
- loop = asyncio.get_event_loop()
- except RuntimeError:
- loop = asyncio.new_event_loop()
- asyncio.set_event_loop(loop)
- if loop.is_closed():
- loop = asyncio.new_event_loop()
- asyncio.set_event_loop(loop)
-
- if loop.is_running():
- result_container = []
-
- def thread_target() -> None:
- nonlocal result_container
- new_loop = asyncio.new_event_loop()
- asyncio.set_event_loop(new_loop)
- try:
- result_container.append(
- new_loop.run_until_complete(self._arun(*args, **kwargs))
- )
- except Exception as e:
- result_container.append(e)
- finally:
- new_loop.close()
-
- thread = threading.Thread(target=thread_target)
- thread.start()
- thread.join()
- result = result_container[0]
- if isinstance(result, Exception):
- raise result
- return result
-
- else:
- result = loop.run_until_complete(self._arun(*args, **kwargs))
- loop.close()
- return result
diff --git a/libs/community/langchain_community/tools/ainetwork/owner.py b/libs/community/langchain_community/tools/ainetwork/owner.py
deleted file mode 100644
index 13d41d9327..0000000000
--- a/libs/community/langchain_community/tools/ainetwork/owner.py
+++ /dev/null
@@ -1,115 +0,0 @@
-import builtins
-import json
-from typing import List, Optional, Type, Union
-
-from langchain_core.callbacks import AsyncCallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.ainetwork.base import AINBaseTool, OperationType
-
-
-class RuleSchema(BaseModel):
- """Schema for owner operations."""
-
- type: OperationType = Field(...)
- path: str = Field(..., description="Blockchain reference path")
- address: Optional[Union[str, List[str]]] = Field(
- None, description="A single address or a list of addresses"
- )
- write_owner: Optional[bool] = Field(
- False, description="Authority to edit the `owner` property of the path"
- )
- write_rule: Optional[bool] = Field(
- False, description="Authority to edit `write rule` for the path"
- )
- write_function: Optional[bool] = Field(
- False, description="Authority to `set function` for the path"
- )
- branch_owner: Optional[bool] = Field(
- False, description="Authority to initialize `owner` of sub-paths"
- )
-
-
-class AINOwnerOps(AINBaseTool):
- """Tool for owner operations."""
-
- name: str = "AINownerOps"
- description: str = """
-Rules for `owner` in AINetwork Blockchain database.
-An address set as `owner` can modify permissions according to its granted authorities
-
-## Path Rule
-- (/[a-zA-Z_0-9]+)+
-- Permission checks ascend from the most specific (child) path to broader (parent) paths until an `owner` is located.
-
-## Address Rules
-- 0x[0-9a-fA-F]{40}: 40-digit hexadecimal address
-- *: All addresses permitted
-- Defaults to the current session's address
-
-## SET
-- `SET` alters permissions for specific addresses, while other addresses remain unaffected.
-- When removing an address of `owner`, set all authorities for that address to false.
-- message `write_owner permission evaluated false` if fail
-
-### Example
-- type: SET
-- path: /apps/langchain
-- address: [, ]
-- write_owner: True
-- write_rule: True
-- write_function: True
-- branch_owner: True
-
-## GET
-- Provides all addresses with `owner` permissions and their authorities in the path.
-
-### Example
-- type: GET
-- path: /apps/langchain
-""" # noqa: E501
- args_schema: Type[BaseModel] = RuleSchema
-
- async def _arun(
- self,
- type: OperationType,
- path: str,
- address: Optional[Union[str, List[str]]] = None,
- write_owner: Optional[bool] = None,
- write_rule: Optional[bool] = None,
- write_function: Optional[bool] = None,
- branch_owner: Optional[bool] = None,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- from ain.types import ValueOnlyTransactionInput
-
- try:
- if type is OperationType.SET:
- if address is None:
- address = self.interface.wallet.defaultAccount.address
- if isinstance(address, str):
- address = [address]
- res = await self.interface.db.ref(path).setOwner(
- transactionInput=ValueOnlyTransactionInput(
- value={
- ".owner": {
- "owners": {
- address: {
- "write_owner": write_owner or False,
- "write_rule": write_rule or False,
- "write_function": write_function or False,
- "branch_owner": branch_owner or False,
- }
- for address in address
- }
- }
- }
- )
- )
- elif type is OperationType.GET:
- res = await self.interface.db.ref(path).getOwner()
- else:
- raise ValueError(f"Unsupported 'type': {type}.")
- return json.dumps(res, ensure_ascii=False)
- except Exception as e:
- return f"{builtins.type(e).__name__}: {str(e)}"
diff --git a/libs/community/langchain_community/tools/ainetwork/rule.py b/libs/community/langchain_community/tools/ainetwork/rule.py
deleted file mode 100644
index 5a24c9e5ab..0000000000
--- a/libs/community/langchain_community/tools/ainetwork/rule.py
+++ /dev/null
@@ -1,82 +0,0 @@
-import builtins
-import json
-from typing import Optional, Type
-
-from langchain_core.callbacks import AsyncCallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.ainetwork.base import AINBaseTool, OperationType
-
-
-class RuleSchema(BaseModel):
- """Schema for owner operations."""
-
- type: OperationType = Field(...)
- path: str = Field(..., description="Path on the blockchain where the rule applies")
- eval: Optional[str] = Field(None, description="eval string to determine permission")
-
-
-class AINRuleOps(AINBaseTool):
- """Tool for owner operations."""
-
- name: str = "AINruleOps"
- description: str = """
-Covers the write `rule` for the AINetwork Blockchain database. The SET type specifies write permissions using the `eval` variable as a JavaScript eval string.
-In order to AINvalueOps with SET at the path, the execution result of the `eval` string must be true.
-
-## Path Rules
-1. Allowed characters for directory: `[a-zA-Z_0-9]`
-2. Use `$` for template variables as directory.
-
-## Eval String Special Variables
-- auth.addr: Address of the writer for the path
-- newData: New data for the path
-- data: Current data for the path
-- currentTime: Time in seconds
-- lastBlockNumber: Latest processed block number
-
-## Eval String Functions
-- getValue()
-- getRule()
-- getOwner()
-- getFunction()
-- evalRule(, , auth, currentTime)
-- evalOwner(, 'write_owner', auth)
-
-## SET Example
-- type: SET
-- path: /apps/langchain_project_1/$from/$to/$img
-- eval: auth.addr===$from&&!getValue('/apps/image_db/'+$img)
-
-## GET Example
-- type: GET
-- path: /apps/langchain_project_1
-""" # noqa: E501
- args_schema: Type[BaseModel] = RuleSchema
-
- async def _arun(
- self,
- type: OperationType,
- path: str,
- eval: Optional[str] = None,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- from ain.types import ValueOnlyTransactionInput
-
- try:
- if type is OperationType.SET:
- if eval is None:
- raise ValueError("'eval' is required for SET operation.")
-
- res = await self.interface.db.ref(path).setRule(
- transactionInput=ValueOnlyTransactionInput(
- value={".rule": {"write": eval}}
- )
- )
- elif type is OperationType.GET:
- res = await self.interface.db.ref(path).getRule()
- else:
- raise ValueError(f"Unsupported 'type': {type}.")
- return json.dumps(res, ensure_ascii=False)
- except Exception as e:
- return f"{builtins.type(e).__name__}: {str(e)}"
diff --git a/libs/community/langchain_community/tools/ainetwork/transfer.py b/libs/community/langchain_community/tools/ainetwork/transfer.py
deleted file mode 100644
index 81d630af2e..0000000000
--- a/libs/community/langchain_community/tools/ainetwork/transfer.py
+++ /dev/null
@@ -1,34 +0,0 @@
-import json
-from typing import Optional, Type
-
-from langchain_core.callbacks import AsyncCallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.ainetwork.base import AINBaseTool
-
-
-class TransferSchema(BaseModel):
- """Schema for transfer operations."""
-
- address: str = Field(..., description="Address to transfer AIN to")
- amount: int = Field(..., description="Amount of AIN to transfer")
-
-
-class AINTransfer(AINBaseTool):
- """Tool for transfer operations."""
-
- name: str = "AINtransfer"
- description: str = "Transfers AIN to a specified address"
- args_schema: Type[TransferSchema] = TransferSchema
-
- async def _arun(
- self,
- address: str,
- amount: int,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- try:
- res = await self.interface.wallet.transfer(address, amount, nonce=-1)
- return json.dumps(res, ensure_ascii=False)
- except Exception as e:
- return f"{type(e).__name__}: {str(e)}"
diff --git a/libs/community/langchain_community/tools/ainetwork/utils.py b/libs/community/langchain_community/tools/ainetwork/utils.py
deleted file mode 100644
index 0bb848d006..0000000000
--- a/libs/community/langchain_community/tools/ainetwork/utils.py
+++ /dev/null
@@ -1,63 +0,0 @@
-"""AINetwork Blockchain tool utils."""
-
-from __future__ import annotations
-
-import os
-from typing import TYPE_CHECKING, Literal, Optional
-
-if TYPE_CHECKING:
- from ain.ain import Ain
-
-
-def authenticate(network: Optional[Literal["mainnet", "testnet"]] = "testnet") -> Ain:
- """Authenticate using the AIN Blockchain"""
-
- try:
- from ain.ain import Ain
- except ImportError as e:
- raise ImportError(
- "Cannot import ain-py related modules. Please install the package with "
- "`pip install ain-py`."
- ) from e
-
- if network == "mainnet":
- provider_url = "https://mainnet-api.ainetwork.ai/"
- chain_id = 1
- if "AIN_BLOCKCHAIN_ACCOUNT_PRIVATE_KEY" in os.environ:
- private_key = os.environ["AIN_BLOCKCHAIN_ACCOUNT_PRIVATE_KEY"]
- else:
- raise EnvironmentError(
- "Error: The AIN_BLOCKCHAIN_ACCOUNT_PRIVATE_KEY environmental variable "
- "has not been set."
- )
- elif network == "testnet":
- provider_url = "https://testnet-api.ainetwork.ai/"
- chain_id = 0
- if "AIN_BLOCKCHAIN_ACCOUNT_PRIVATE_KEY" in os.environ:
- private_key = os.environ["AIN_BLOCKCHAIN_ACCOUNT_PRIVATE_KEY"]
- else:
- raise EnvironmentError(
- "Error: The AIN_BLOCKCHAIN_ACCOUNT_PRIVATE_KEY environmental variable "
- "has not been set."
- )
- elif network is None:
- if (
- "AIN_BLOCKCHAIN_PROVIDER_URL" in os.environ
- and "AIN_BLOCKCHAIN_CHAIN_ID" in os.environ
- and "AIN_BLOCKCHAIN_ACCOUNT_PRIVATE_KEY" in os.environ
- ):
- provider_url = os.environ["AIN_BLOCKCHAIN_PROVIDER_URL"]
- chain_id = int(os.environ["AIN_BLOCKCHAIN_CHAIN_ID"])
- private_key = os.environ["AIN_BLOCKCHAIN_ACCOUNT_PRIVATE_KEY"]
- else:
- raise EnvironmentError(
- "Error: The AIN_BLOCKCHAIN_PROVIDER_URL and "
- "AIN_BLOCKCHAIN_ACCOUNT_PRIVATE_KEY and AIN_BLOCKCHAIN_CHAIN_ID "
- "environmental variable has not been set."
- )
- else:
- raise ValueError(f"Unsupported 'network': {network}")
-
- ain = Ain(provider_url, chain_id)
- ain.wallet.addAndSetDefaultAccount(private_key)
- return ain
diff --git a/libs/community/langchain_community/tools/ainetwork/value.py b/libs/community/langchain_community/tools/ainetwork/value.py
deleted file mode 100644
index be6414727f..0000000000
--- a/libs/community/langchain_community/tools/ainetwork/value.py
+++ /dev/null
@@ -1,85 +0,0 @@
-import builtins
-import json
-from typing import Optional, Type, Union
-
-from langchain_core.callbacks import AsyncCallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.ainetwork.base import AINBaseTool, OperationType
-
-
-class ValueSchema(BaseModel):
- """Schema for value operations."""
-
- type: OperationType = Field(...)
- path: str = Field(..., description="Blockchain reference path")
- value: Optional[Union[int, str, float, dict]] = Field(
- None, description="Value to be set at the path"
- )
-
-
-class AINValueOps(AINBaseTool):
- """Tool for value operations."""
-
- name: str = "AINvalueOps"
- description: str = """
-Covers the read and write value for the AINetwork Blockchain database.
-
-## SET
-- Set a value at a given path
-
-### Example
-- type: SET
-- path: /apps/langchain_test_1/object
-- value: {1: 2, "34": 56}
-
-## GET
-- Retrieve a value at a given path
-
-### Example
-- type: GET
-- path: /apps/langchain_test_1/DB
-
-## Special paths
-- `/accounts//balance`: Account balance
-- `/accounts//nonce`: Account nonce
-- `/apps`: Applications
-- `/consensus`: Consensus
-- `/checkin`: Check-in
-- `/deposit///`: Deposit
-- `/deposit_accounts///`: Deposit accounts
-- `/escrow`: Escrow
-- `/payments`: Payment
-- `/sharding`: Sharding
-- `/token/name`: Token name
-- `/token/symbol`: Token symbol
-- `/token/total_supply`: Token total supply
-- `/transfer////value`: Transfer
-- `/withdraw///`: Withdraw
-"""
- args_schema: Type[BaseModel] = ValueSchema
-
- async def _arun(
- self,
- type: OperationType,
- path: str,
- value: Optional[Union[int, str, float, dict]] = None,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- from ain.types import ValueOnlyTransactionInput
-
- try:
- if type is OperationType.SET:
- if value is None:
- raise ValueError("'value' is required for SET operation.")
-
- res = await self.interface.db.ref(path).setValue(
- transactionInput=ValueOnlyTransactionInput(value=value)
- )
- elif type is OperationType.GET:
- res = await self.interface.db.ref(path).getValue()
- else:
- raise ValueError(f"Unsupported 'type': {type}.")
- return json.dumps(res, ensure_ascii=False)
- except Exception as e:
- return f"{builtins.type(e).__name__}: {str(e)}"
diff --git a/libs/community/langchain_community/tools/amadeus/__init__.py b/libs/community/langchain_community/tools/amadeus/__init__.py
deleted file mode 100644
index 570958f809..0000000000
--- a/libs/community/langchain_community/tools/amadeus/__init__.py
+++ /dev/null
@@ -1,9 +0,0 @@
-"""Amadeus tools."""
-
-from langchain_community.tools.amadeus.closest_airport import AmadeusClosestAirport
-from langchain_community.tools.amadeus.flight_search import AmadeusFlightSearch
-
-__all__ = [
- "AmadeusClosestAirport",
- "AmadeusFlightSearch",
-]
diff --git a/libs/community/langchain_community/tools/amadeus/base.py b/libs/community/langchain_community/tools/amadeus/base.py
deleted file mode 100644
index 3fd3f377ce..0000000000
--- a/libs/community/langchain_community/tools/amadeus/base.py
+++ /dev/null
@@ -1,19 +0,0 @@
-"""Base class for Amadeus tools."""
-
-from __future__ import annotations
-
-from typing import TYPE_CHECKING
-
-from langchain_core.tools import BaseTool
-from pydantic import Field
-
-from langchain_community.tools.amadeus.utils import authenticate
-
-if TYPE_CHECKING:
- from amadeus import Client
-
-
-class AmadeusBaseTool(BaseTool):
- """Base Tool for Amadeus."""
-
- client: Client = Field(default_factory=authenticate)
diff --git a/libs/community/langchain_community/tools/amadeus/closest_airport.py b/libs/community/langchain_community/tools/amadeus/closest_airport.py
deleted file mode 100644
index 9523f73afb..0000000000
--- a/libs/community/langchain_community/tools/amadeus/closest_airport.py
+++ /dev/null
@@ -1,62 +0,0 @@
-from typing import Any, Dict, Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.language_models import BaseLanguageModel
-from pydantic import BaseModel, Field, model_validator
-
-from langchain_community.chat_models import ChatOpenAI
-from langchain_community.tools.amadeus.base import AmadeusBaseTool
-
-
-class ClosestAirportSchema(BaseModel):
- """Schema for the AmadeusClosestAirport tool."""
-
- location: str = Field(
- description=(
- " The location for which you would like to find the nearest airport "
- " along with optional details such as country, state, region, or "
- " province, allowing for easy processing and identification of "
- " the closest airport. Examples of the format are the following:\n"
- " Cali, Colombia\n "
- " Lincoln, Nebraska, United States\n"
- " New York, United States\n"
- " Sydney, New South Wales, Australia\n"
- " Rome, Lazio, Italy\n"
- " Toronto, Ontario, Canada\n"
- )
- )
-
-
-class AmadeusClosestAirport(AmadeusBaseTool):
- """Tool for finding the closest airport to a particular location."""
-
- name: str = "closest_airport"
- description: str = (
- "Use this tool to find the closest airport to a particular location."
- )
- args_schema: Type[ClosestAirportSchema] = ClosestAirportSchema
-
- llm: Optional[BaseLanguageModel] = Field(default=None)
- """Tool's llm used for calculating the closest airport. Defaults to `ChatOpenAI`."""
-
- @model_validator(mode="before")
- @classmethod
- def set_llm(cls, values: Dict[str, Any]) -> Any:
- if not values.get("llm"):
- # For backward-compatibility
- values["llm"] = ChatOpenAI(temperature=0)
- return values
-
- def _run(
- self,
- location: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- content = (
- f" What is the nearest airport to {location}? Please respond with the "
- " airport's International Air Transport Association (IATA) Location "
- ' Identifier in the following JSON format. JSON: "iataCode": "IATA '
- ' Location Identifier" '
- )
-
- return self.llm.invoke(content) # type: ignore[union-attr]
diff --git a/libs/community/langchain_community/tools/amadeus/flight_search.py b/libs/community/langchain_community/tools/amadeus/flight_search.py
deleted file mode 100644
index c3cd8fe7bb..0000000000
--- a/libs/community/langchain_community/tools/amadeus/flight_search.py
+++ /dev/null
@@ -1,153 +0,0 @@
-import logging
-from datetime import datetime as dt
-from typing import Dict, Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.amadeus.base import AmadeusBaseTool
-
-logger = logging.getLogger(__name__)
-
-
-class FlightSearchSchema(BaseModel):
- """Schema for the AmadeusFlightSearch tool."""
-
- originLocationCode: str = Field(
- description=(
- " The three letter International Air Transport "
- " Association (IATA) Location Identifier for the "
- " search's origin airport. "
- )
- )
- destinationLocationCode: str = Field(
- description=(
- " The three letter International Air Transport "
- " Association (IATA) Location Identifier for the "
- " search's destination airport. "
- )
- )
- departureDateTimeEarliest: str = Field(
- description=(
- " The earliest departure datetime from the origin airport "
- " for the flight search in the following format: "
- ' "YYYY-MM-DDTHH:MM:SS", where "T" separates the date and time '
- ' components. For example: "2023-06-09T10:30:00" represents '
- " June 9th, 2023, at 10:30 AM. "
- )
- )
- departureDateTimeLatest: str = Field(
- description=(
- " The latest departure datetime from the origin airport "
- " for the flight search in the following format: "
- ' "YYYY-MM-DDTHH:MM:SS", where "T" separates the date and time '
- ' components. For example: "2023-06-09T10:30:00" represents '
- " June 9th, 2023, at 10:30 AM. "
- )
- )
- page_number: int = Field(
- default=1,
- description="The specific page number of flight results to retrieve",
- )
-
-
-class AmadeusFlightSearch(AmadeusBaseTool):
- """Tool for searching for a single flight between two airports."""
-
- name: str = "single_flight_search"
- description: str = (
- " Use this tool to search for a single flight between the origin and "
- " destination airports at a departure between an earliest and "
- " latest datetime. "
- )
- args_schema: Type[FlightSearchSchema] = FlightSearchSchema
-
- def _run(
- self,
- originLocationCode: str,
- destinationLocationCode: str,
- departureDateTimeEarliest: str,
- departureDateTimeLatest: str,
- page_number: int = 1,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> list:
- try:
- from amadeus import ResponseError
- except ImportError as e:
- raise ImportError(
- "Unable to import amadeus, please install with `pip install amadeus`."
- ) from e
-
- RESULTS_PER_PAGE = 10
-
- # Authenticate and retrieve a client
- client = self.client
-
- # Check that earliest and latest dates are in the same day
- earliestDeparture = dt.strptime(departureDateTimeEarliest, "%Y-%m-%dT%H:%M:%S")
- latestDeparture = dt.strptime(departureDateTimeLatest, "%Y-%m-%dT%H:%M:%S")
-
- if earliestDeparture.date() != latestDeparture.date():
- logger.error(
- " Error: Earliest and latest departure dates need to be the "
- " same date. If you're trying to search for round-trip "
- " flights, call this function for the outbound flight first, "
- " and then call again for the return flight. "
- )
- return [None]
-
- # Collect all results from the Amadeus Flight Offers Search API
- response = None
- try:
- response = client.shopping.flight_offers_search.get(
- originLocationCode=originLocationCode,
- destinationLocationCode=destinationLocationCode,
- departureDate=latestDeparture.strftime("%Y-%m-%d"),
- adults=1,
- )
- except ResponseError as error:
- print(error) # noqa: T201
-
- # Generate output dictionary
- output = []
- if response is not None:
- for offer in response.data:
- itinerary: Dict = {}
- itinerary["price"] = {}
- itinerary["price"]["total"] = offer["price"]["total"]
- currency = offer["price"]["currency"]
- currency = response.result["dictionaries"]["currencies"][currency]
- itinerary["price"]["currency"] = {}
- itinerary["price"]["currency"] = currency
-
- segments = []
- for segment in offer["itineraries"][0]["segments"]:
- flight = {}
- flight["departure"] = segment["departure"]
- flight["arrival"] = segment["arrival"]
- flight["flightNumber"] = segment["number"]
- carrier = segment["carrierCode"]
- carrier = response.result["dictionaries"]["carriers"][carrier]
- flight["carrier"] = carrier
-
- segments.append(flight)
-
- itinerary["segments"] = []
- itinerary["segments"] = segments
-
- output.append(itinerary)
-
- # Filter out flights after latest departure time
- for index, offer in enumerate(output):
- offerDeparture = dt.strptime(
- offer["segments"][0]["departure"]["at"], "%Y-%m-%dT%H:%M:%S"
- )
-
- if offerDeparture > latestDeparture:
- output.pop(index)
-
- # Return the paginated results
- startIndex = (page_number - 1) * RESULTS_PER_PAGE
- endIndex = startIndex + RESULTS_PER_PAGE
-
- return output[startIndex:endIndex]
diff --git a/libs/community/langchain_community/tools/amadeus/utils.py b/libs/community/langchain_community/tools/amadeus/utils.py
deleted file mode 100644
index 7fef81c237..0000000000
--- a/libs/community/langchain_community/tools/amadeus/utils.py
+++ /dev/null
@@ -1,43 +0,0 @@
-"""O365 tool utils."""
-
-from __future__ import annotations
-
-import logging
-import os
-from typing import TYPE_CHECKING
-
-if TYPE_CHECKING:
- from amadeus import Client
-
-logger = logging.getLogger(__name__)
-
-
-def authenticate() -> Client:
- """Authenticate using the Amadeus API"""
- try:
- from amadeus import Client
- except ImportError as e:
- raise ImportError(
- "Cannot import amadeus. Please install the package with "
- "`pip install amadeus`."
- ) from e
-
- if "AMADEUS_CLIENT_ID" in os.environ and "AMADEUS_CLIENT_SECRET" in os.environ:
- client_id = os.environ["AMADEUS_CLIENT_ID"]
- client_secret = os.environ["AMADEUS_CLIENT_SECRET"]
- else:
- logger.error(
- "Error: The AMADEUS_CLIENT_ID and AMADEUS_CLIENT_SECRET environmental "
- "variables have not been set. Visit the following link on how to "
- "acquire these authorization tokens: "
- "https://developers.amadeus.com/register"
- )
- return None
-
- hostname = "test" # Default hostname
- if "AMADEUS_HOSTNAME" in os.environ:
- hostname = os.environ["AMADEUS_HOSTNAME"]
-
- client = Client(client_id=client_id, client_secret=client_secret, hostname=hostname)
-
- return client
diff --git a/libs/community/langchain_community/tools/arxiv/__init__.py b/libs/community/langchain_community/tools/arxiv/__init__.py
deleted file mode 100644
index a240b3f2ea..0000000000
--- a/libs/community/langchain_community/tools/arxiv/__init__.py
+++ /dev/null
@@ -1,6 +0,0 @@
-from langchain_community.tools.arxiv.tool import ArxivQueryRun
-
-"""Arxiv API toolkit."""
-"""Tool for the Arxiv Search API."""
-
-__all__ = ["ArxivQueryRun"]
diff --git a/libs/community/langchain_community/tools/arxiv/tool.py b/libs/community/langchain_community/tools/arxiv/tool.py
deleted file mode 100644
index 601022a132..0000000000
--- a/libs/community/langchain_community/tools/arxiv/tool.py
+++ /dev/null
@@ -1,39 +0,0 @@
-"""Tool for the Arxiv API."""
-
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.utilities.arxiv import ArxivAPIWrapper
-
-
-class ArxivInput(BaseModel):
- """Input for the Arxiv tool."""
-
- query: str = Field(description="search query to look up")
-
-
-class ArxivQueryRun(BaseTool):
- """Tool that searches the Arxiv API."""
-
- name: str = "arxiv"
- description: str = (
- "A wrapper around Arxiv.org "
- "Useful for when you need to answer questions about Physics, Mathematics, "
- "Computer Science, Quantitative Biology, Quantitative Finance, Statistics, "
- "Electrical Engineering, and Economics "
- "from scientific articles on arxiv.org. "
- "Input should be a search query."
- )
- api_wrapper: ArxivAPIWrapper = Field(default_factory=ArxivAPIWrapper) # type: ignore[arg-type]
- args_schema: Type[BaseModel] = ArxivInput
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Arxiv tool."""
- return self.api_wrapper.run(query)
diff --git a/libs/community/langchain_community/tools/asknews/__init__.py b/libs/community/langchain_community/tools/asknews/__init__.py
deleted file mode 100644
index 635745a7d7..0000000000
--- a/libs/community/langchain_community/tools/asknews/__init__.py
+++ /dev/null
@@ -1,7 +0,0 @@
-"""AskNews API toolkit."""
-
-from langchain_community.tools.asknews.tool import (
- AskNewsSearch,
-)
-
-__all__ = ["AskNewsSearch"]
diff --git a/libs/community/langchain_community/tools/asknews/tool.py b/libs/community/langchain_community/tools/asknews/tool.py
deleted file mode 100644
index ca5de6970c..0000000000
--- a/libs/community/langchain_community/tools/asknews/tool.py
+++ /dev/null
@@ -1,82 +0,0 @@
-"""
-Tool for the AskNews API.
-
-To use this tool, you must first set your credentials as environment variables:
- ASKNEWS_CLIENT_ID
- ASKNEWS_CLIENT_SECRET
-"""
-
-from typing import Any, Optional, Type
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.utilities.asknews import AskNewsAPIWrapper
-
-
-class SearchInput(BaseModel):
- """Input for the AskNews Search tool."""
-
- query: str = Field(
- description="Search query to be used for finding real-time or historical news "
- "information."
- )
- hours_back: Optional[int] = Field(
- 0,
- description="If the Assistant deems that the event may have occurred more "
- "than 48 hours ago, it estimates the number of hours back to search. For "
- "example, if the event was one month ago, the Assistant may set this to 720. "
- "One week would be 168. The Assistant can estimate up to on year back (8760).",
- )
-
-
-class AskNewsSearch(BaseTool):
- """Tool that searches the AskNews API."""
-
- name: str = "asknews_search"
- description: str = (
- "This tool allows you to perform a search on up-to-date news and historical "
- "news. If you needs news from more than 48 hours ago, you can estimate the "
- "number of hours back to search."
- )
- api_wrapper: AskNewsAPIWrapper = Field(default_factory=AskNewsAPIWrapper)
- max_results: int = 10
- args_schema: Optional[Type[BaseModel]] = SearchInput
-
- def _run(
- self,
- query: str,
- hours_back: int = 0,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- **kwargs: Any,
- ) -> str:
- """Use the tool."""
- try:
- return self.api_wrapper.search_news(
- query,
- hours_back=hours_back,
- max_results=self.max_results,
- )
- except Exception as e:
- return repr(e)
-
- async def _arun(
- self,
- query: str,
- hours_back: int = 0,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- **kwargs: Any,
- ) -> str:
- """Use the tool asynchronously."""
- try:
- return await self.api_wrapper.asearch_news(
- query,
- hours_back=hours_back,
- max_results=self.max_results,
- )
- except Exception as e:
- return repr(e)
diff --git a/libs/community/langchain_community/tools/audio/__init__.py b/libs/community/langchain_community/tools/audio/__init__.py
deleted file mode 100644
index 9024dc6fea..0000000000
--- a/libs/community/langchain_community/tools/audio/__init__.py
+++ /dev/null
@@ -1,7 +0,0 @@
-from langchain_community.tools.audio.huggingface_text_to_speech_inference import (
- HuggingFaceTextToSpeechModelInference,
-)
-
-__all__ = [
- "HuggingFaceTextToSpeechModelInference",
-]
diff --git a/libs/community/langchain_community/tools/audio/huggingface_text_to_speech_inference.py b/libs/community/langchain_community/tools/audio/huggingface_text_to_speech_inference.py
deleted file mode 100644
index 9e20870224..0000000000
--- a/libs/community/langchain_community/tools/audio/huggingface_text_to_speech_inference.py
+++ /dev/null
@@ -1,122 +0,0 @@
-import logging
-import os
-import uuid
-from datetime import datetime
-from typing import Callable, Literal, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import SecretStr
-
-logger = logging.getLogger(__name__)
-
-
-class HuggingFaceTextToSpeechModelInference(BaseTool):
- """HuggingFace Text-to-Speech Model Inference.
-
- Requirements:
- - Environment variable ``HUGGINGFACE_API_KEY`` must be set,
- or passed as a named parameter to the constructor.
- """
-
- name: str = "openai_text_to_speech"
- """Name of the tool."""
- description: str = "A wrapper around OpenAI Text-to-Speech API. "
- """Description of the tool."""
-
- model: str
- """Model name."""
- file_extension: str
- """File extension of the output audio file."""
- destination_dir: str
- """Directory to save the output audio file."""
- file_namer: Callable[[], str]
- """Function to generate unique file names."""
-
- api_url: str
- huggingface_api_key: SecretStr
-
- _HUGGINGFACE_API_KEY_ENV_NAME: str = "HUGGINGFACE_API_KEY"
- _HUGGINGFACE_API_URL_ROOT: str = "https://api-inference.huggingface.co/models"
-
- def __init__(
- self,
- model: str,
- file_extension: str,
- *,
- destination_dir: str = "./tts",
- file_naming_func: Literal["uuid", "timestamp"] = "uuid",
- huggingface_api_key: Optional[SecretStr] = None,
- _HUGGINGFACE_API_KEY_ENV_NAME: str = "HUGGINGFACE_API_KEY",
- _HUGGINGFACE_API_URL_ROOT: str = "https://api-inference.huggingface.co/models",
- ) -> None:
- if not huggingface_api_key:
- huggingface_api_key = SecretStr(
- os.getenv(_HUGGINGFACE_API_KEY_ENV_NAME, "")
- )
-
- if (
- not huggingface_api_key
- or not huggingface_api_key.get_secret_value()
- or huggingface_api_key.get_secret_value() == ""
- ):
- raise ValueError(
- f"'{_HUGGINGFACE_API_KEY_ENV_NAME}' must be or set or passed"
- )
-
- if file_naming_func == "uuid":
- file_namer = lambda: str(uuid.uuid4()) # noqa: E731
- elif file_naming_func == "timestamp":
- file_namer = lambda: str(int(datetime.now().timestamp())) # noqa: E731
- else:
- raise ValueError(
- f"Invalid value for 'file_naming_func': {file_naming_func}"
- )
-
- super().__init__(
- model=model,
- file_extension=file_extension,
- api_url=f"{_HUGGINGFACE_API_URL_ROOT}/{model}",
- destination_dir=destination_dir,
- file_namer=file_namer,
- huggingface_api_key=huggingface_api_key,
- _HUGGINGFACE_API_KEY_ENV_NAME=_HUGGINGFACE_API_KEY_ENV_NAME,
- _HUGGINGFACE_API_URL_ROOT=_HUGGINGFACE_API_URL_ROOT,
- )
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- response = requests.post(
- self.api_url,
- headers={
- "Authorization": f"Bearer {self.huggingface_api_key.get_secret_value()}"
- },
- json={"inputs": query},
- )
- audio_bytes = response.content
-
- try:
- os.makedirs(self.destination_dir, exist_ok=True)
- except Exception as e:
- logger.error(f"Error creating directory '{self.destination_dir}': {e}")
- raise
-
- output_file = os.path.join(
- self.destination_dir,
- f"{str(self.file_namer())}.{self.file_extension}",
- )
-
- try:
- with open(output_file, mode="xb") as f:
- f.write(audio_bytes)
- except FileExistsError:
- raise ValueError("Output name must be unique")
- except Exception as e:
- logger.error(f"Error occurred while creating file: {e}")
- raise
-
- return output_file
diff --git a/libs/community/langchain_community/tools/azure_ai_services/__init__.py b/libs/community/langchain_community/tools/azure_ai_services/__init__.py
deleted file mode 100644
index 637285ac63..0000000000
--- a/libs/community/langchain_community/tools/azure_ai_services/__init__.py
+++ /dev/null
@@ -1,25 +0,0 @@
-"""Azure AI Services Tools."""
-
-from langchain_community.tools.azure_ai_services.document_intelligence import (
- AzureAiServicesDocumentIntelligenceTool,
-)
-from langchain_community.tools.azure_ai_services.image_analysis import (
- AzureAiServicesImageAnalysisTool,
-)
-from langchain_community.tools.azure_ai_services.speech_to_text import (
- AzureAiServicesSpeechToTextTool,
-)
-from langchain_community.tools.azure_ai_services.text_analytics_for_health import (
- AzureAiServicesTextAnalyticsForHealthTool,
-)
-from langchain_community.tools.azure_ai_services.text_to_speech import (
- AzureAiServicesTextToSpeechTool,
-)
-
-__all__ = [
- "AzureAiServicesDocumentIntelligenceTool",
- "AzureAiServicesImageAnalysisTool",
- "AzureAiServicesSpeechToTextTool",
- "AzureAiServicesTextToSpeechTool",
- "AzureAiServicesTextAnalyticsForHealthTool",
-]
diff --git a/libs/community/langchain_community/tools/azure_ai_services/document_intelligence.py b/libs/community/langchain_community/tools/azure_ai_services/document_intelligence.py
deleted file mode 100644
index cd0ac25018..0000000000
--- a/libs/community/langchain_community/tools/azure_ai_services/document_intelligence.py
+++ /dev/null
@@ -1,146 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Dict, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import model_validator
-
-from langchain_community.tools.azure_ai_services.utils import (
- detect_file_src_type,
-)
-
-logger = logging.getLogger(__name__)
-
-
-class AzureAiServicesDocumentIntelligenceTool(BaseTool):
- """Tool that queries the Azure AI Services Document Intelligence API.
-
- In order to set this up, follow instructions at:
- https://learn.microsoft.com/en-us/azure/ai-services/document-intelligence/quickstarts/get-started-sdks-rest-api?view=doc-intel-4.0.0&pivots=programming-language-python
- """
-
- azure_ai_services_key: str = "" #: :meta private:
- azure_ai_services_endpoint: str = "" #: :meta private:
- doc_analysis_client: Any #: :meta private:
-
- name: str = "azure_ai_services_document_intelligence"
- description: str = (
- "A wrapper around Azure AI Services Document Intelligence. "
- "Useful for when you need to "
- "extract text, tables, and key-value pairs from documents. "
- "Input should be a url to a document."
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and endpoint exists in environment."""
- azure_ai_services_key = get_from_dict_or_env(
- values, "azure_ai_services_key", "AZURE_AI_SERVICES_KEY"
- )
-
- azure_ai_services_endpoint = get_from_dict_or_env(
- values, "azure_ai_services_endpoint", "AZURE_AI_SERVICES_ENDPOINT"
- )
-
- try:
- from azure.ai.formrecognizer import DocumentAnalysisClient
- from azure.core.credentials import AzureKeyCredential
-
- values["doc_analysis_client"] = DocumentAnalysisClient(
- endpoint=azure_ai_services_endpoint,
- credential=AzureKeyCredential(azure_ai_services_key),
- )
-
- except ImportError:
- raise ImportError(
- "azure-ai-formrecognizer is not installed. "
- "Run `pip install azure-ai-formrecognizer` to install."
- )
-
- return values
-
- def _parse_tables(self, tables: List[Any]) -> List[Any]:
- result = []
- for table in tables:
- rc, cc = table.row_count, table.column_count
- _table = [["" for _ in range(cc)] for _ in range(rc)]
- for cell in table.cells:
- _table[cell.row_index][cell.column_index] = cell.content
- result.append(_table)
- return result
-
- def _parse_kv_pairs(self, kv_pairs: List[Any]) -> List[Any]:
- result = []
- for kv_pair in kv_pairs:
- key = kv_pair.key.content if kv_pair.key else ""
- value = kv_pair.value.content if kv_pair.value else ""
- result.append((key, value))
- return result
-
- def _document_analysis(self, document_path: str) -> Dict:
- document_src_type = detect_file_src_type(document_path)
- if document_src_type == "local":
- with open(document_path, "rb") as document:
- poller = self.doc_analysis_client.begin_analyze_document(
- "prebuilt-document", document
- )
- elif document_src_type == "remote":
- poller = self.doc_analysis_client.begin_analyze_document_from_url(
- "prebuilt-document", document_path
- )
- else:
- raise ValueError(f"Invalid document path: {document_path}")
-
- result = poller.result()
- res_dict = {}
-
- if result.content is not None:
- res_dict["content"] = result.content
-
- if result.tables is not None:
- res_dict["tables"] = self._parse_tables(result.tables)
-
- if result.key_value_pairs is not None:
- res_dict["key_value_pairs"] = self._parse_kv_pairs(result.key_value_pairs)
-
- return res_dict
-
- def _format_document_analysis_result(self, document_analysis_result: Dict) -> str:
- formatted_result = []
- if "content" in document_analysis_result:
- formatted_result.append(
- f"Content: {document_analysis_result['content']}".replace("\n", " ")
- )
-
- if "tables" in document_analysis_result:
- for i, table in enumerate(document_analysis_result["tables"]):
- formatted_result.append(f"Table {i}: {table}".replace("\n", " "))
-
- if "key_value_pairs" in document_analysis_result:
- for kv_pair in document_analysis_result["key_value_pairs"]:
- formatted_result.append(
- f"{kv_pair[0]}: {kv_pair[1]}".replace("\n", " ")
- )
-
- return "\n".join(formatted_result)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- try:
- document_analysis_result = self._document_analysis(query)
- if not document_analysis_result:
- return "No good document analysis result was found"
-
- return self._format_document_analysis_result(document_analysis_result)
- except Exception as e:
- raise RuntimeError(
- f"Error while running AzureAiServicesDocumentIntelligenceTool: {e}"
- )
diff --git a/libs/community/langchain_community/tools/azure_ai_services/image_analysis.py b/libs/community/langchain_community/tools/azure_ai_services/image_analysis.py
deleted file mode 100644
index c3292cff6b..0000000000
--- a/libs/community/langchain_community/tools/azure_ai_services/image_analysis.py
+++ /dev/null
@@ -1,195 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Dict, Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import model_validator
-
-from langchain_community.tools.azure_ai_services.utils import (
- detect_file_src_type,
-)
-
-logger = logging.getLogger(__name__)
-
-
-class AzureAiServicesImageAnalysisTool(BaseTool):
- """Tool that queries the Azure AI Services Image Analysis API.
-
- In order to set this up, follow instructions at:
- https://learn.microsoft.com/azure/ai-services/computer-vision/quickstarts-sdk/image-analysis-client-library-40
-
- Attributes:
- azure_ai_services_key (Optional[str]): The API key for Azure AI Services.
- azure_ai_services_endpoint (Optional[str]): The endpoint URL for Azure AI Services.
- visual_features Any: The visual features to analyze in the image, can be set as
- either strings or azure.ai.vision.imageanalysis.models.VisualFeatures.
- (e.g. 'TAGS', VisualFeatures.CAPTION).
- image_analysis_client (Any): The client for interacting
- with Azure AI Services Image Analysis.
- name (str): The name of the tool.
- description (str): A description of the tool,
- including its purpose and expected input.
- """
-
- azure_ai_services_key: Optional[str] = None #: :meta private:
- azure_ai_services_endpoint: Optional[str] = None #: :meta private:
- visual_features: Any = None
- image_analysis_client: Any = None #: :meta private:
-
- name: str = "azure_ai_services_image_analysis"
- description: str = (
- "A wrapper around Azure AI Services Image Analysis. "
- "Useful for when you need to analyze images. "
- "Input must be a url string or path string to an image."
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and endpoint exists in environment."""
- azure_ai_services_key = get_from_dict_or_env(
- values, "azure_ai_services_key", "AZURE_AI_SERVICES_KEY"
- )
-
- azure_ai_services_endpoint = get_from_dict_or_env(
- values, "azure_ai_services_endpoint", "AZURE_AI_SERVICES_ENDPOINT"
- )
-
- """Validate that azure-ai-vision-imageanalysis is installed."""
- try:
- from azure.ai.vision.imageanalysis import ImageAnalysisClient
- from azure.ai.vision.imageanalysis.models import VisualFeatures
- from azure.core.credentials import AzureKeyCredential
- except ImportError:
- raise ImportError(
- "azure-ai-vision-imageanalysis is not installed. "
- "Run `pip install azure-ai-vision-imageanalysis` to install. "
- )
-
- """Validate Azure AI Vision Image Analysis client can be initialized."""
- try:
- values["image_analysis_client"] = ImageAnalysisClient(
- endpoint=azure_ai_services_endpoint,
- credential=AzureKeyCredential(azure_ai_services_key),
- )
- except Exception as e:
- raise RuntimeError(
- f"Initialization of Azure AI Vision Image Analysis client failed: {e}"
- )
-
- visual_features = values.get(
- "visual_features",
- [
- VisualFeatures.TAGS,
- VisualFeatures.OBJECTS,
- VisualFeatures.CAPTION,
- VisualFeatures.READ,
- ],
- )
- values["visual_features"] = visual_features
- return values
-
- def _image_analysis(self, image_path: str) -> Dict:
- try:
- from azure.ai.vision.imageanalysis import ImageAnalysisClient
- except ImportError:
- pass
-
- self.image_analysis_client: ImageAnalysisClient
-
- image_src_type = detect_file_src_type(image_path)
- if image_src_type == "local":
- with open(image_path, "rb") as image_file:
- image_data = image_file.read()
- result = self.image_analysis_client.analyze(
- image_data=image_data,
- visual_features=self.visual_features,
- )
- elif image_src_type == "remote":
- result = self.image_analysis_client.analyze_from_url(
- image_url=image_path,
- visual_features=self.visual_features,
- )
- else:
- raise ValueError(f"Invalid image path: {image_path}")
-
- res_dict = {}
- if result:
- if result.caption is not None:
- res_dict["caption"] = result.caption.text
-
- if result.objects is not None:
- res_dict["objects"] = [obj.tags[0].name for obj in result.objects.list]
-
- if result.tags is not None:
- res_dict["tags"] = [tag.name for tag in result.tags.list]
-
- if result.read is not None and len(result.read.blocks) > 0:
- res_dict["text"] = [line.text for line in result.read.blocks[0].lines]
-
- if result.dense_captions is not None and len(result.dense_captions) > 0:
- res_dict["dense_captions"] = [
- str(dc) for dc in result.dense_captions.list
- ]
-
- if result.smart_crops is not None and len(result.smart_crops) > 0:
- res_dict["smart_crops"] = [str(sc) for sc in result.smart_crops.list]
-
- if result.people is not None and len(result.people) > 0:
- res_dict["people"] = [str(p) for p in result.people.list]
-
- return res_dict
-
- def _format_image_analysis_result(self, image_analysis_result: Dict) -> str:
- formatted_result = []
- if "caption" in image_analysis_result:
- formatted_result.append("Caption: " + image_analysis_result["caption"])
-
- if (
- "objects" in image_analysis_result
- and len(image_analysis_result["objects"]) > 0
- ):
- formatted_result.append(
- "Objects: " + ", ".join(image_analysis_result["objects"])
- )
-
- if "tags" in image_analysis_result and len(image_analysis_result["tags"]) > 0:
- formatted_result.append("Tags: " + ", ".join(image_analysis_result["tags"]))
-
- if "text" in image_analysis_result and len(image_analysis_result["text"]) > 0:
- formatted_result.append("Text: " + ", ".join(image_analysis_result["text"]))
-
- if "dense_captions" in image_analysis_result:
- formatted_result.append(
- "Dense Captions: " + ", ".join(image_analysis_result["dense_captions"])
- )
-
- if "smart_crops" in image_analysis_result:
- formatted_result.append(
- "Smart Crops: " + ", ".join(image_analysis_result["smart_crops"])
- )
-
- if "people" in image_analysis_result:
- formatted_result.append(
- "People: " + ", ".join(image_analysis_result["people"])
- )
-
- return "\n".join(formatted_result)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- try:
- image_analysis_result = self._image_analysis(query)
- if not image_analysis_result:
- return "No good image analysis result was found"
-
- return self._format_image_analysis_result(image_analysis_result)
- except Exception as e:
- raise RuntimeError(f"Error while running AzureAiImageAnalysisTool: {e}")
diff --git a/libs/community/langchain_community/tools/azure_ai_services/speech_to_text.py b/libs/community/langchain_community/tools/azure_ai_services/speech_to_text.py
deleted file mode 100644
index 15e08d2722..0000000000
--- a/libs/community/langchain_community/tools/azure_ai_services/speech_to_text.py
+++ /dev/null
@@ -1,123 +0,0 @@
-from __future__ import annotations
-
-import logging
-import time
-from typing import Any, Dict, Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import model_validator
-
-from langchain_community.tools.azure_ai_services.utils import (
- detect_file_src_type,
- download_audio_from_url,
-)
-
-logger = logging.getLogger(__name__)
-
-
-class AzureAiServicesSpeechToTextTool(BaseTool):
- """Tool that queries the Azure AI Services Speech to Text API.
-
- In order to set this up, follow instructions at:
- https://learn.microsoft.com/en-us/azure/ai-services/speech-service/get-started-speech-to-text?pivots=programming-language-python
- """
-
- azure_ai_services_key: str = "" #: :meta private:
- azure_ai_services_region: str = "" #: :meta private:
- speech_language: str = "en-US" #: :meta private:
- speech_config: Any #: :meta private:
-
- name: str = "azure_ai_services_speech_to_text"
- description: str = (
- "A wrapper around Azure AI Services Speech to Text. "
- "Useful for when you need to transcribe audio to text. "
- "Input should be a url to an audio file."
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and endpoint exists in environment."""
- azure_ai_services_key = get_from_dict_or_env(
- values, "azure_ai_services_key", "AZURE_AI_SERVICES_KEY"
- )
-
- azure_ai_services_region = get_from_dict_or_env(
- values, "azure_ai_services_region", "AZURE_AI_SERVICES_REGION"
- )
-
- try:
- import azure.cognitiveservices.speech as speechsdk
-
- values["speech_config"] = speechsdk.SpeechConfig(
- subscription=azure_ai_services_key, region=azure_ai_services_region
- )
- except ImportError:
- raise ImportError(
- "azure-cognitiveservices-speech is not installed. "
- "Run `pip install azure-cognitiveservices-speech` to install."
- )
-
- return values
-
- def _continuous_recognize(self, speech_recognizer: Any) -> str:
- done = False
- text = ""
-
- def stop_cb(evt: Any) -> None:
- """callback that stop continuous recognition"""
- speech_recognizer.stop_continuous_recognition_async()
- nonlocal done
- done = True
-
- def retrieve_cb(evt: Any) -> None:
- """callback that retrieves the intermediate recognition results"""
- nonlocal text
- text += evt.result.text
-
- # retrieve text on recognized events
- speech_recognizer.recognized.connect(retrieve_cb)
- # stop continuous recognition on either session stopped or canceled events
- speech_recognizer.session_stopped.connect(stop_cb)
- speech_recognizer.canceled.connect(stop_cb)
-
- # Start continuous speech recognition
- speech_recognizer.start_continuous_recognition_async()
- while not done:
- time.sleep(0.5)
- return text
-
- def _speech_to_text(self, audio_path: str, speech_language: str) -> str:
- try:
- import azure.cognitiveservices.speech as speechsdk
- except ImportError:
- pass
-
- audio_src_type = detect_file_src_type(audio_path)
- if audio_src_type == "local":
- audio_config = speechsdk.AudioConfig(filename=audio_path)
- elif audio_src_type == "remote":
- tmp_audio_path = download_audio_from_url(audio_path)
- audio_config = speechsdk.AudioConfig(filename=tmp_audio_path)
- else:
- raise ValueError(f"Invalid audio path: {audio_path}")
-
- self.speech_config.speech_recognition_language = speech_language
- speech_recognizer = speechsdk.SpeechRecognizer(self.speech_config, audio_config)
- return self._continuous_recognize(speech_recognizer)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- try:
- text = self._speech_to_text(query, self.speech_language)
- return text
- except Exception as e:
- raise RuntimeError(
- f"Error while running AzureAiServicesSpeechToTextTool: {e}"
- )
diff --git a/libs/community/langchain_community/tools/azure_ai_services/text_analytics_for_health.py b/libs/community/langchain_community/tools/azure_ai_services/text_analytics_for_health.py
deleted file mode 100644
index 6df15788f5..0000000000
--- a/libs/community/langchain_community/tools/azure_ai_services/text_analytics_for_health.py
+++ /dev/null
@@ -1,105 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Dict, Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import model_validator
-
-logger = logging.getLogger(__name__)
-
-
-class AzureAiServicesTextAnalyticsForHealthTool(BaseTool):
- """Tool that queries the Azure AI Services Text Analytics for Health API.
-
- In order to set this up, follow instructions at:
- https://learn.microsoft.com/en-us/azure/ai-services/language-service/text-analytics-for-health/quickstart?pivots=programming-language-python
- """
-
- azure_ai_services_key: str = "" #: :meta private:
- azure_ai_services_endpoint: str = "" #: :meta private:
- text_analytics_client: Any #: :meta private:
-
- name: str = "azure_ai_services_text_analytics_for_health"
- description: str = (
- "A wrapper around Azure AI Services Text Analytics for Health. "
- "Useful for when you need to identify entities in healthcare data. "
- "Input should be text."
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and endpoint exists in environment."""
- azure_ai_services_key = get_from_dict_or_env(
- values, "azure_ai_services_key", "AZURE_AI_SERVICES_KEY"
- )
-
- azure_ai_services_endpoint = get_from_dict_or_env(
- values, "azure_ai_services_endpoint", "AZURE_AI_SERVICES_ENDPOINT"
- )
-
- try:
- import azure.ai.textanalytics as sdk
- from azure.core.credentials import AzureKeyCredential
-
- values["text_analytics_client"] = sdk.TextAnalyticsClient(
- endpoint=azure_ai_services_endpoint,
- credential=AzureKeyCredential(azure_ai_services_key),
- )
-
- except ImportError:
- raise ImportError(
- "azure-ai-textanalytics is not installed. "
- "Run `pip install azure-ai-textanalytics` to install."
- )
-
- return values
-
- def _text_analysis(self, text: str) -> Dict:
- poller = self.text_analytics_client.begin_analyze_healthcare_entities(
- [{"id": "1", "language": "en", "text": text}]
- )
-
- result = poller.result()
-
- res_dict = {}
-
- docs = [doc for doc in result if not doc.is_error]
-
- if docs is not None:
- res_dict["entities"] = [
- f"{x.text} is a healthcare entity of type {x.category}"
- for y in docs
- for x in y.entities
- ]
-
- return res_dict
-
- def _format_text_analysis_result(self, text_analysis_result: Dict) -> str:
- formatted_result = []
- if "entities" in text_analysis_result:
- formatted_result.append(
- f"""The text contains the following healthcare entities: {
- ", ".join(text_analysis_result["entities"])
- }""".replace("\n", " ")
- )
-
- return "\n".join(formatted_result)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- try:
- text_analysis_result = self._text_analysis(query)
-
- return self._format_text_analysis_result(text_analysis_result)
- except Exception as e:
- raise RuntimeError(
- f"Error while running AzureAiServicesTextAnalyticsForHealthTool: {e}"
- )
diff --git a/libs/community/langchain_community/tools/azure_ai_services/text_to_speech.py b/libs/community/langchain_community/tools/azure_ai_services/text_to_speech.py
deleted file mode 100644
index 1291e2dac5..0000000000
--- a/libs/community/langchain_community/tools/azure_ai_services/text_to_speech.py
+++ /dev/null
@@ -1,106 +0,0 @@
-from __future__ import annotations
-
-import logging
-import tempfile
-from typing import Any, Dict, Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import model_validator
-
-logger = logging.getLogger(__name__)
-
-
-class AzureAiServicesTextToSpeechTool(BaseTool):
- """Tool that queries the Azure AI Services Text to Speech API.
-
- In order to set this up, follow instructions at:
- https://learn.microsoft.com/en-us/azure/ai-services/speech-service/get-started-text-to-speech?pivots=programming-language-python
- """
-
- name: str = "azure_ai_services_text_to_speech"
- description: str = (
- "A wrapper around Azure AI Services Text to Speech API. "
- "Useful for when you need to convert text to speech. "
- )
- return_direct: bool = True
-
- azure_ai_services_key: str = "" #: :meta private:
- azure_ai_services_region: str = "" #: :meta private:
- speech_language: str = "en-US" #: :meta private:
- speech_config: Any #: :meta private:
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and endpoint exists in environment."""
- azure_ai_services_key = get_from_dict_or_env(
- values, "azure_ai_services_key", "AZURE_AI_SERVICES_KEY"
- )
-
- azure_ai_services_region = get_from_dict_or_env(
- values, "azure_ai_services_region", "AZURE_AI_SERVICES_REGION"
- )
-
- try:
- import azure.cognitiveservices.speech as speechsdk
-
- values["speech_config"] = speechsdk.SpeechConfig(
- subscription=azure_ai_services_key, region=azure_ai_services_region
- )
- except ImportError:
- raise ImportError(
- "azure-cognitiveservices-speech is not installed. "
- "Run `pip install azure-cognitiveservices-speech` to install."
- )
-
- return values
-
- def _text_to_speech(self, text: str, speech_language: str) -> str:
- try:
- import azure.cognitiveservices.speech as speechsdk
- except ImportError:
- pass
-
- self.speech_config.speech_synthesis_language = speech_language
- speech_synthesizer = speechsdk.SpeechSynthesizer(
- speech_config=self.speech_config, audio_config=None
- )
- result = speech_synthesizer.speak_text(text)
-
- if result.reason == speechsdk.ResultReason.SynthesizingAudioCompleted:
- stream = speechsdk.AudioDataStream(result)
- with tempfile.NamedTemporaryFile(
- mode="wb", suffix=".wav", delete=False
- ) as f:
- stream.save_to_wav_file(f.name)
-
- return f.name
-
- elif result.reason == speechsdk.ResultReason.Canceled:
- cancellation_details = result.cancellation_details
- logger.debug(f"Speech synthesis canceled: {cancellation_details.reason}")
- if cancellation_details.reason == speechsdk.CancellationReason.Error:
- raise RuntimeError(
- f"Speech synthesis error: {cancellation_details.error_details}"
- )
-
- return "Speech synthesis canceled."
-
- else:
- return f"Speech synthesis failed: {result.reason}"
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- try:
- speech_file = self._text_to_speech(query, self.speech_language)
- return speech_file
- except Exception as e:
- raise RuntimeError(
- f"Error while running AzureAiServicesTextToSpeechTool: {e}"
- )
diff --git a/libs/community/langchain_community/tools/azure_ai_services/utils.py b/libs/community/langchain_community/tools/azure_ai_services/utils.py
deleted file mode 100644
index 9de8f923b7..0000000000
--- a/libs/community/langchain_community/tools/azure_ai_services/utils.py
+++ /dev/null
@@ -1,29 +0,0 @@
-import os
-import tempfile
-from urllib.parse import urlparse
-
-import requests
-
-
-def detect_file_src_type(file_path: str) -> str:
- """Detect if the file is local or remote."""
- if os.path.isfile(file_path):
- return "local"
-
- parsed_url = urlparse(file_path)
- if parsed_url.scheme and parsed_url.netloc:
- return "remote"
-
- return "invalid"
-
-
-def download_audio_from_url(audio_url: str) -> str:
- """Download audio from url to local."""
- ext = audio_url.split(".")[-1]
- response = requests.get(audio_url, stream=True)
- response.raise_for_status()
- with tempfile.NamedTemporaryFile(mode="wb", suffix=f".{ext}", delete=False) as f:
- for chunk in response.iter_content(chunk_size=8192):
- f.write(chunk)
-
- return f.name
diff --git a/libs/community/langchain_community/tools/azure_cognitive_services/__init__.py b/libs/community/langchain_community/tools/azure_cognitive_services/__init__.py
deleted file mode 100644
index 1121e4e89d..0000000000
--- a/libs/community/langchain_community/tools/azure_cognitive_services/__init__.py
+++ /dev/null
@@ -1,25 +0,0 @@
-"""Azure Cognitive Services Tools."""
-
-from langchain_community.tools.azure_cognitive_services.form_recognizer import (
- AzureCogsFormRecognizerTool,
-)
-from langchain_community.tools.azure_cognitive_services.image_analysis import (
- AzureCogsImageAnalysisTool,
-)
-from langchain_community.tools.azure_cognitive_services.speech2text import (
- AzureCogsSpeech2TextTool,
-)
-from langchain_community.tools.azure_cognitive_services.text2speech import (
- AzureCogsText2SpeechTool,
-)
-from langchain_community.tools.azure_cognitive_services.text_analytics_health import (
- AzureCogsTextAnalyticsHealthTool,
-)
-
-__all__ = [
- "AzureCogsImageAnalysisTool",
- "AzureCogsFormRecognizerTool",
- "AzureCogsSpeech2TextTool",
- "AzureCogsText2SpeechTool",
- "AzureCogsTextAnalyticsHealthTool",
-]
diff --git a/libs/community/langchain_community/tools/azure_cognitive_services/form_recognizer.py b/libs/community/langchain_community/tools/azure_cognitive_services/form_recognizer.py
deleted file mode 100644
index 937b1fc793..0000000000
--- a/libs/community/langchain_community/tools/azure_cognitive_services/form_recognizer.py
+++ /dev/null
@@ -1,144 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Dict, List, Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import model_validator
-
-from langchain_community.tools.azure_cognitive_services.utils import (
- detect_file_src_type,
-)
-
-logger = logging.getLogger(__name__)
-
-
-class AzureCogsFormRecognizerTool(BaseTool):
- """Tool that queries the Azure Cognitive Services Form Recognizer API.
-
- In order to set this up, follow instructions at:
- https://learn.microsoft.com/en-us/azure/applied-ai-services/form-recognizer/quickstarts/get-started-sdks-rest-api?view=form-recog-3.0.0&pivots=programming-language-python
- """
-
- azure_cogs_key: str = "" #: :meta private:
- azure_cogs_endpoint: str = "" #: :meta private:
- doc_analysis_client: Any #: :meta private:
-
- name: str = "azure_cognitive_services_form_recognizer"
- description: str = (
- "A wrapper around Azure Cognitive Services Form Recognizer. "
- "Useful for when you need to "
- "extract text, tables, and key-value pairs from documents. "
- "Input should be a url to a document."
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and endpoint exists in environment."""
- azure_cogs_key = get_from_dict_or_env(
- values, "azure_cogs_key", "AZURE_COGS_KEY"
- )
-
- azure_cogs_endpoint = get_from_dict_or_env(
- values, "azure_cogs_endpoint", "AZURE_COGS_ENDPOINT"
- )
-
- try:
- from azure.ai.formrecognizer import DocumentAnalysisClient
- from azure.core.credentials import AzureKeyCredential
-
- values["doc_analysis_client"] = DocumentAnalysisClient(
- endpoint=azure_cogs_endpoint,
- credential=AzureKeyCredential(azure_cogs_key),
- )
-
- except ImportError:
- raise ImportError(
- "azure-ai-formrecognizer is not installed. "
- "Run `pip install azure-ai-formrecognizer` to install."
- )
-
- return values
-
- def _parse_tables(self, tables: List[Any]) -> List[Any]:
- result = []
- for table in tables:
- rc, cc = table.row_count, table.column_count
- _table = [["" for _ in range(cc)] for _ in range(rc)]
- for cell in table.cells:
- _table[cell.row_index][cell.column_index] = cell.content
- result.append(_table)
- return result
-
- def _parse_kv_pairs(self, kv_pairs: List[Any]) -> List[Any]:
- result = []
- for kv_pair in kv_pairs:
- key = kv_pair.key.content if kv_pair.key else ""
- value = kv_pair.value.content if kv_pair.value else ""
- result.append((key, value))
- return result
-
- def _document_analysis(self, document_path: str) -> Dict:
- document_src_type = detect_file_src_type(document_path)
- if document_src_type == "local":
- with open(document_path, "rb") as document:
- poller = self.doc_analysis_client.begin_analyze_document(
- "prebuilt-document", document
- )
- elif document_src_type == "remote":
- poller = self.doc_analysis_client.begin_analyze_document_from_url(
- "prebuilt-document", document_path
- )
- else:
- raise ValueError(f"Invalid document path: {document_path}")
-
- result = poller.result()
- res_dict = {}
-
- if result.content is not None:
- res_dict["content"] = result.content
-
- if result.tables is not None:
- res_dict["tables"] = self._parse_tables(result.tables)
-
- if result.key_value_pairs is not None:
- res_dict["key_value_pairs"] = self._parse_kv_pairs(result.key_value_pairs)
-
- return res_dict
-
- def _format_document_analysis_result(self, document_analysis_result: Dict) -> str:
- formatted_result = []
- if "content" in document_analysis_result:
- formatted_result.append(
- f"Content: {document_analysis_result['content']}".replace("\n", " ")
- )
-
- if "tables" in document_analysis_result:
- for i, table in enumerate(document_analysis_result["tables"]):
- formatted_result.append(f"Table {i}: {table}".replace("\n", " "))
-
- if "key_value_pairs" in document_analysis_result:
- for kv_pair in document_analysis_result["key_value_pairs"]:
- formatted_result.append(
- f"{kv_pair[0]}: {kv_pair[1]}".replace("\n", " ")
- )
-
- return "\n".join(formatted_result)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- try:
- document_analysis_result = self._document_analysis(query)
- if not document_analysis_result:
- return "No good document analysis result was found"
-
- return self._format_document_analysis_result(document_analysis_result)
- except Exception as e:
- raise RuntimeError(f"Error while running AzureCogsFormRecognizerTool: {e}")
diff --git a/libs/community/langchain_community/tools/azure_cognitive_services/image_analysis.py b/libs/community/langchain_community/tools/azure_cognitive_services/image_analysis.py
deleted file mode 100644
index ce076243a7..0000000000
--- a/libs/community/langchain_community/tools/azure_cognitive_services/image_analysis.py
+++ /dev/null
@@ -1,148 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Dict, Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import model_validator
-
-from langchain_community.tools.azure_cognitive_services.utils import (
- detect_file_src_type,
-)
-
-logger = logging.getLogger(__name__)
-
-
-class AzureCogsImageAnalysisTool(BaseTool):
- """Tool that queries the Azure Cognitive Services Image Analysis API.
-
- In order to set this up, follow instructions at:
- https://learn.microsoft.com/en-us/azure/cognitive-services/computer-vision/quickstarts-sdk/image-analysis-client-library-40
- """
-
- azure_cogs_key: str = "" #: :meta private:
- azure_cogs_endpoint: str = "" #: :meta private:
- vision_service: Any #: :meta private:
- analysis_options: Any #: :meta private:
-
- name: str = "azure_cognitive_services_image_analysis"
- description: str = (
- "A wrapper around Azure Cognitive Services Image Analysis. "
- "Useful for when you need to analyze images. "
- "Input should be a url to an image."
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and endpoint exists in environment."""
- azure_cogs_key = get_from_dict_or_env(
- values, "azure_cogs_key", "AZURE_COGS_KEY"
- )
-
- azure_cogs_endpoint = get_from_dict_or_env(
- values, "azure_cogs_endpoint", "AZURE_COGS_ENDPOINT"
- )
-
- try:
- import azure.ai.vision as sdk
-
- values["vision_service"] = sdk.VisionServiceOptions(
- endpoint=azure_cogs_endpoint, key=azure_cogs_key
- )
-
- values["analysis_options"] = sdk.ImageAnalysisOptions()
- values["analysis_options"].features = (
- sdk.ImageAnalysisFeature.CAPTION
- | sdk.ImageAnalysisFeature.OBJECTS
- | sdk.ImageAnalysisFeature.TAGS
- | sdk.ImageAnalysisFeature.TEXT
- )
- except ImportError:
- raise ImportError(
- "azure-ai-vision is not installed. "
- "Run `pip install azure-ai-vision` to install."
- )
-
- return values
-
- def _image_analysis(self, image_path: str) -> Dict:
- try:
- import azure.ai.vision as sdk
- except ImportError:
- pass
-
- image_src_type = detect_file_src_type(image_path)
- if image_src_type == "local":
- vision_source = sdk.VisionSource(filename=image_path)
- elif image_src_type == "remote":
- vision_source = sdk.VisionSource(url=image_path)
- else:
- raise ValueError(f"Invalid image path: {image_path}")
-
- image_analyzer = sdk.ImageAnalyzer(
- self.vision_service, vision_source, self.analysis_options
- )
- result = image_analyzer.analyze()
-
- res_dict = {}
- if result.reason == sdk.ImageAnalysisResultReason.ANALYZED:
- if result.caption is not None:
- res_dict["caption"] = result.caption.content
-
- if result.objects is not None:
- res_dict["objects"] = [obj.name for obj in result.objects]
-
- if result.tags is not None:
- res_dict["tags"] = [tag.name for tag in result.tags]
-
- if result.text is not None:
- res_dict["text"] = [line.content for line in result.text.lines]
-
- else:
- error_details = sdk.ImageAnalysisErrorDetails.from_result(result)
- raise RuntimeError(
- f"Image analysis failed.\n"
- f"Reason: {error_details.reason}\n"
- f"Details: {error_details.message}"
- )
-
- return res_dict
-
- def _format_image_analysis_result(self, image_analysis_result: Dict) -> str:
- formatted_result = []
- if "caption" in image_analysis_result:
- formatted_result.append("Caption: " + image_analysis_result["caption"])
-
- if (
- "objects" in image_analysis_result
- and len(image_analysis_result["objects"]) > 0
- ):
- formatted_result.append(
- "Objects: " + ", ".join(image_analysis_result["objects"])
- )
-
- if "tags" in image_analysis_result and len(image_analysis_result["tags"]) > 0:
- formatted_result.append("Tags: " + ", ".join(image_analysis_result["tags"]))
-
- if "text" in image_analysis_result and len(image_analysis_result["text"]) > 0:
- formatted_result.append("Text: " + ", ".join(image_analysis_result["text"]))
-
- return "\n".join(formatted_result)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- try:
- image_analysis_result = self._image_analysis(query)
- if not image_analysis_result:
- return "No good image analysis result was found"
-
- return self._format_image_analysis_result(image_analysis_result)
- except Exception as e:
- raise RuntimeError(f"Error while running AzureCogsImageAnalysisTool: {e}")
diff --git a/libs/community/langchain_community/tools/azure_cognitive_services/speech2text.py b/libs/community/langchain_community/tools/azure_cognitive_services/speech2text.py
deleted file mode 100644
index 125c910df1..0000000000
--- a/libs/community/langchain_community/tools/azure_cognitive_services/speech2text.py
+++ /dev/null
@@ -1,121 +0,0 @@
-from __future__ import annotations
-
-import logging
-import time
-from typing import Any, Dict, Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import model_validator
-
-from langchain_community.tools.azure_cognitive_services.utils import (
- detect_file_src_type,
- download_audio_from_url,
-)
-
-logger = logging.getLogger(__name__)
-
-
-class AzureCogsSpeech2TextTool(BaseTool):
- """Tool that queries the Azure Cognitive Services Speech2Text API.
-
- In order to set this up, follow instructions at:
- https://learn.microsoft.com/en-us/azure/cognitive-services/speech-service/get-started-speech-to-text?pivots=programming-language-python
- """
-
- azure_cogs_key: str = "" #: :meta private:
- azure_cogs_region: str = "" #: :meta private:
- speech_language: str = "en-US" #: :meta private:
- speech_config: Any #: :meta private:
-
- name: str = "azure_cognitive_services_speech2text"
- description: str = (
- "A wrapper around Azure Cognitive Services Speech2Text. "
- "Useful for when you need to transcribe audio to text. "
- "Input should be a url to an audio file."
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and endpoint exists in environment."""
- azure_cogs_key = get_from_dict_or_env(
- values, "azure_cogs_key", "AZURE_COGS_KEY"
- )
-
- azure_cogs_region = get_from_dict_or_env(
- values, "azure_cogs_region", "AZURE_COGS_REGION"
- )
-
- try:
- import azure.cognitiveservices.speech as speechsdk
-
- values["speech_config"] = speechsdk.SpeechConfig(
- subscription=azure_cogs_key, region=azure_cogs_region
- )
- except ImportError:
- raise ImportError(
- "azure-cognitiveservices-speech is not installed. "
- "Run `pip install azure-cognitiveservices-speech` to install."
- )
-
- return values
-
- def _continuous_recognize(self, speech_recognizer: Any) -> str:
- done = False
- text = ""
-
- def stop_cb(evt: Any) -> None:
- """callback that stop continuous recognition"""
- speech_recognizer.stop_continuous_recognition_async()
- nonlocal done
- done = True
-
- def retrieve_cb(evt: Any) -> None:
- """callback that retrieves the intermediate recognition results"""
- nonlocal text
- text += evt.result.text
-
- # retrieve text on recognized events
- speech_recognizer.recognized.connect(retrieve_cb)
- # stop continuous recognition on either session stopped or canceled events
- speech_recognizer.session_stopped.connect(stop_cb)
- speech_recognizer.canceled.connect(stop_cb)
-
- # Start continuous speech recognition
- speech_recognizer.start_continuous_recognition_async()
- while not done:
- time.sleep(0.5)
- return text
-
- def _speech2text(self, audio_path: str, speech_language: str) -> str:
- try:
- import azure.cognitiveservices.speech as speechsdk
- except ImportError:
- pass
-
- audio_src_type = detect_file_src_type(audio_path)
- if audio_src_type == "local":
- audio_config = speechsdk.AudioConfig(filename=audio_path)
- elif audio_src_type == "remote":
- tmp_audio_path = download_audio_from_url(audio_path)
- audio_config = speechsdk.AudioConfig(filename=tmp_audio_path)
- else:
- raise ValueError(f"Invalid audio path: {audio_path}")
-
- self.speech_config.speech_recognition_language = speech_language
- speech_recognizer = speechsdk.SpeechRecognizer(self.speech_config, audio_config)
- return self._continuous_recognize(speech_recognizer)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- try:
- text = self._speech2text(query, self.speech_language)
- return text
- except Exception as e:
- raise RuntimeError(f"Error while running AzureCogsSpeech2TextTool: {e}")
diff --git a/libs/community/langchain_community/tools/azure_cognitive_services/text2speech.py b/libs/community/langchain_community/tools/azure_cognitive_services/text2speech.py
deleted file mode 100644
index 343653fe9c..0000000000
--- a/libs/community/langchain_community/tools/azure_cognitive_services/text2speech.py
+++ /dev/null
@@ -1,103 +0,0 @@
-from __future__ import annotations
-
-import logging
-import tempfile
-from typing import Any, Dict, Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import model_validator
-
-logger = logging.getLogger(__name__)
-
-
-class AzureCogsText2SpeechTool(BaseTool):
- """Tool that queries the Azure Cognitive Services Text2Speech API.
-
- In order to set this up, follow instructions at:
- https://learn.microsoft.com/en-us/azure/cognitive-services/speech-service/get-started-text-to-speech?pivots=programming-language-python
- """
-
- azure_cogs_key: str = "" #: :meta private:
- azure_cogs_region: str = "" #: :meta private:
- speech_language: str = "en-US" #: :meta private:
- speech_config: Any #: :meta private:
-
- name: str = "azure_cognitive_services_text2speech"
- description: str = (
- "A wrapper around Azure Cognitive Services Text2Speech. "
- "Useful for when you need to convert text to speech. "
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and endpoint exists in environment."""
- azure_cogs_key = get_from_dict_or_env(
- values, "azure_cogs_key", "AZURE_COGS_KEY"
- )
-
- azure_cogs_region = get_from_dict_or_env(
- values, "azure_cogs_region", "AZURE_COGS_REGION"
- )
-
- try:
- import azure.cognitiveservices.speech as speechsdk
-
- values["speech_config"] = speechsdk.SpeechConfig(
- subscription=azure_cogs_key, region=azure_cogs_region
- )
- except ImportError:
- raise ImportError(
- "azure-cognitiveservices-speech is not installed. "
- "Run `pip install azure-cognitiveservices-speech` to install."
- )
-
- return values
-
- def _text2speech(self, text: str, speech_language: str) -> str:
- try:
- import azure.cognitiveservices.speech as speechsdk
- except ImportError:
- pass
-
- self.speech_config.speech_synthesis_language = speech_language
- speech_synthesizer = speechsdk.SpeechSynthesizer(
- speech_config=self.speech_config, audio_config=None
- )
- result = speech_synthesizer.speak_text(text)
-
- if result.reason == speechsdk.ResultReason.SynthesizingAudioCompleted:
- stream = speechsdk.AudioDataStream(result)
- with tempfile.NamedTemporaryFile(
- mode="wb", suffix=".wav", delete=False
- ) as f:
- stream.save_to_wav_file(f.name)
-
- return f.name
-
- elif result.reason == speechsdk.ResultReason.Canceled:
- cancellation_details = result.cancellation_details
- logger.debug(f"Speech synthesis canceled: {cancellation_details.reason}")
- if cancellation_details.reason == speechsdk.CancellationReason.Error:
- raise RuntimeError(
- f"Speech synthesis error: {cancellation_details.error_details}"
- )
-
- return "Speech synthesis canceled."
-
- else:
- return f"Speech synthesis failed: {result.reason}"
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- try:
- speech_file = self._text2speech(query, self.speech_language)
- return speech_file
- except Exception as e:
- raise RuntimeError(f"Error while running AzureCogsText2SpeechTool: {e}")
diff --git a/libs/community/langchain_community/tools/azure_cognitive_services/text_analytics_health.py b/libs/community/langchain_community/tools/azure_cognitive_services/text_analytics_health.py
deleted file mode 100644
index 26864a8382..0000000000
--- a/libs/community/langchain_community/tools/azure_cognitive_services/text_analytics_health.py
+++ /dev/null
@@ -1,105 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Dict, Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import model_validator
-
-logger = logging.getLogger(__name__)
-
-
-class AzureCogsTextAnalyticsHealthTool(BaseTool):
- """Tool that queries the Azure Cognitive Services Text Analytics for Health API.
-
- In order to set this up, follow instructions at:
- https://learn.microsoft.com/en-us/azure/ai-services/language-service/text-analytics-for-health/quickstart?tabs=windows&pivots=programming-language-python
- """
-
- azure_cogs_key: str = "" #: :meta private:
- azure_cogs_endpoint: str = "" #: :meta private:
- text_analytics_client: Any #: :meta private:
-
- name: str = "azure_cognitive_services_text_analyics_health"
- description: str = (
- "A wrapper around Azure Cognitive Services Text Analytics for Health. "
- "Useful for when you need to identify entities in healthcare data. "
- "Input should be text."
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key and endpoint exists in environment."""
- azure_cogs_key = get_from_dict_or_env(
- values, "azure_cogs_key", "AZURE_COGS_KEY"
- )
-
- azure_cogs_endpoint = get_from_dict_or_env(
- values, "azure_cogs_endpoint", "AZURE_COGS_ENDPOINT"
- )
-
- try:
- import azure.ai.textanalytics as sdk
- from azure.core.credentials import AzureKeyCredential
-
- values["text_analytics_client"] = sdk.TextAnalyticsClient(
- endpoint=azure_cogs_endpoint,
- credential=AzureKeyCredential(azure_cogs_key),
- )
-
- except ImportError:
- raise ImportError(
- "azure-ai-textanalytics is not installed. "
- "Run `pip install azure-ai-textanalytics` to install."
- )
-
- return values
-
- def _text_analysis(self, text: str) -> Dict:
- poller = self.text_analytics_client.begin_analyze_healthcare_entities(
- [{"id": "1", "language": "en", "text": text}]
- )
-
- result = poller.result()
-
- res_dict = {}
-
- docs = [doc for doc in result if not doc.is_error]
-
- if docs is not None:
- res_dict["entities"] = [
- f"{x.text} is a healthcare entity of type {x.category}"
- for y in docs
- for x in y.entities
- ]
-
- return res_dict
-
- def _format_text_analysis_result(self, text_analysis_result: Dict) -> str:
- formatted_result = []
- if "entities" in text_analysis_result:
- formatted_result.append(
- f"""The text contains the following healthcare entities: {
- ", ".join(text_analysis_result["entities"])
- }""".replace("\n", " ")
- )
-
- return "\n".join(formatted_result)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- try:
- text_analysis_result = self._text_analysis(query)
-
- return self._format_text_analysis_result(text_analysis_result)
- except Exception as e:
- raise RuntimeError(
- f"Error while running AzureCogsTextAnalyticsHealthTool: {e}"
- )
diff --git a/libs/community/langchain_community/tools/azure_cognitive_services/utils.py b/libs/community/langchain_community/tools/azure_cognitive_services/utils.py
deleted file mode 100644
index 9de8f923b7..0000000000
--- a/libs/community/langchain_community/tools/azure_cognitive_services/utils.py
+++ /dev/null
@@ -1,29 +0,0 @@
-import os
-import tempfile
-from urllib.parse import urlparse
-
-import requests
-
-
-def detect_file_src_type(file_path: str) -> str:
- """Detect if the file is local or remote."""
- if os.path.isfile(file_path):
- return "local"
-
- parsed_url = urlparse(file_path)
- if parsed_url.scheme and parsed_url.netloc:
- return "remote"
-
- return "invalid"
-
-
-def download_audio_from_url(audio_url: str) -> str:
- """Download audio from url to local."""
- ext = audio_url.split(".")[-1]
- response = requests.get(audio_url, stream=True)
- response.raise_for_status()
- with tempfile.NamedTemporaryFile(mode="wb", suffix=f".{ext}", delete=False) as f:
- for chunk in response.iter_content(chunk_size=8192):
- f.write(chunk)
-
- return f.name
diff --git a/libs/community/langchain_community/tools/bearly/__init__.py b/libs/community/langchain_community/tools/bearly/__init__.py
deleted file mode 100644
index e69de29bb2..0000000000
diff --git a/libs/community/langchain_community/tools/bearly/tool.py b/libs/community/langchain_community/tools/bearly/tool.py
deleted file mode 100644
index eba71c0d7f..0000000000
--- a/libs/community/langchain_community/tools/bearly/tool.py
+++ /dev/null
@@ -1,165 +0,0 @@
-import base64
-import itertools
-import json
-import re
-from pathlib import Path
-from typing import Dict, List, Type
-
-import requests
-from langchain_core.tools import Tool
-from pydantic import BaseModel, Field
-
-
-def strip_markdown_code(md_string: str) -> str:
- """Strip markdown code from a string."""
- stripped_string = re.sub(r"^`{1,3}.*?\n", "", md_string, flags=re.DOTALL)
- stripped_string = re.sub(r"`{1,3}$", "", stripped_string)
- return stripped_string
-
-
-def head_file(path: str, n: int) -> List[str]:
- """Get the first n lines of a file."""
- try:
- with open(path, "r") as f:
- return [str(line) for line in itertools.islice(f, n)]
- except Exception:
- return []
-
-
-def file_to_base64(path: str) -> str:
- """Convert a file to base64."""
- with open(path, "rb") as f:
- return base64.b64encode(f.read()).decode()
-
-
-class BearlyInterpreterToolArguments(BaseModel):
- """Arguments for the BearlyInterpreterTool."""
-
- python_code: str = Field(
- ...,
- examples=["print('Hello World')"],
- description=(
- "The pure python script to be evaluated. "
- "The contents will be in main.py. "
- "It should not be in markdown format."
- ),
- )
-
-
-base_description = """Evaluates python code in a sandbox environment. \
-The environment resets on every execution. \
-You must send the whole script every time and print your outputs. \
-Script should be pure python code that can be evaluated. \
-It should be in python format NOT markdown. \
-The code should NOT be wrapped in backticks. \
-All python packages including requests, matplotlib, scipy, numpy, pandas, \
-etc are available. \
-If you have any files outputted write them to "output/" relative to the execution \
-path. Output can only be read from the directory, stdout, and stdin. \
-Do not use things like plot.show() as it will \
-not work instead write them out `output/` and a link to the file will be returned. \
-print() any output and results so you can capture the output."""
-
-
-class FileInfo(BaseModel):
- """Information about a file to be uploaded."""
-
- source_path: str
- description: str
- target_path: str
-
-
-class BearlyInterpreterTool:
- """Tool for evaluating python code in a sandbox environment."""
-
- api_key: str
- endpoint: str = "https://exec.bearly.ai/v1/interpreter"
- name: str = "bearly_interpreter"
- args_schema: Type[BaseModel] = BearlyInterpreterToolArguments
- files: Dict[str, FileInfo] = {}
-
- def __init__(self, api_key: str):
- self.api_key = api_key
-
- @property
- def file_description(self) -> str:
- if len(self.files) == 0:
- return ""
- lines = ["The following files available in the evaluation environment:"]
- for target_path, file_info in self.files.items():
- peek_content = head_file(file_info.source_path, 4)
- lines.append(
- f"- path: `{target_path}` \n first four lines: {peek_content}"
- f" \n description: `{file_info.description}`"
- )
- return "\n".join(lines)
-
- @property
- def description(self) -> str:
- return (base_description + "\n\n" + self.file_description).strip()
-
- def make_input_files(self) -> List[dict]:
- files = []
- for target_path, file_info in self.files.items():
- files.append(
- {
- "pathname": target_path,
- "contentsBasesixtyfour": file_to_base64(file_info.source_path),
- }
- )
- return files
-
- def _run(self, python_code: str) -> dict:
- script = strip_markdown_code(python_code)
- resp = requests.post(
- "https://exec.bearly.ai/v1/interpreter",
- data=json.dumps(
- {
- "fileContents": script,
- "inputFiles": self.make_input_files(),
- "outputDir": "output/",
- "outputAsLinks": True,
- }
- ),
- headers={"Authorization": self.api_key},
- ).json()
- return {
- "stdout": (
- base64.b64decode(resp["stdoutBasesixtyfour"]).decode()
- if resp["stdoutBasesixtyfour"]
- else ""
- ),
- "stderr": (
- base64.b64decode(resp["stderrBasesixtyfour"]).decode()
- if resp["stderrBasesixtyfour"]
- else ""
- ),
- "fileLinks": resp["fileLinks"],
- "exitCode": resp["exitCode"],
- }
-
- async def _arun(self, query: str) -> str:
- """Use the tool asynchronously."""
- raise NotImplementedError("custom_search does not support async")
-
- def add_file(self, source_path: str, target_path: str, description: str) -> None:
- if target_path in self.files:
- raise ValueError("target_path already exists")
- if not Path(source_path).exists():
- raise ValueError("source_path does not exist")
- self.files[target_path] = FileInfo(
- target_path=target_path, source_path=source_path, description=description
- )
-
- def clear_files(self) -> None:
- self.files = {}
-
- # TODO: this is because we can't have a dynamic description
- # because of the base pydantic class
- def as_tool(self) -> Tool:
- return Tool.from_function(
- func=self._run,
- name=self.name,
- description=self.description,
- args_schema=self.args_schema,
- )
diff --git a/libs/community/langchain_community/tools/bing_search/__init__.py b/libs/community/langchain_community/tools/bing_search/__init__.py
deleted file mode 100644
index b5e133a05a..0000000000
--- a/libs/community/langchain_community/tools/bing_search/__init__.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""Bing Search API toolkit."""
-
-from langchain_community.tools.bing_search.tool import BingSearchResults, BingSearchRun
-
-__all__ = ["BingSearchRun", "BingSearchResults"]
diff --git a/libs/community/langchain_community/tools/bing_search/tool.py b/libs/community/langchain_community/tools/bing_search/tool.py
deleted file mode 100644
index 9c05405f82..0000000000
--- a/libs/community/langchain_community/tools/bing_search/tool.py
+++ /dev/null
@@ -1,98 +0,0 @@
-"""Tool for the Bing search API."""
-
-from typing import Dict, List, Literal, Optional, Tuple
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-
-from langchain_community.utilities.bing_search import BingSearchAPIWrapper
-
-
-class BingSearchRun(BaseTool):
- """Tool that queries the Bing search API."""
-
- name: str = "bing_search"
- description: str = (
- "A wrapper around Bing Search. "
- "Useful for when you need to answer questions about current events. "
- "Input should be a search query."
- )
- api_wrapper: BingSearchAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return self.api_wrapper.run(query)
-
-
-class BingSearchResults(BaseTool):
- """Bing Search tool.
-
- Setup:
- Install ``langchain-community`` and set environment variable ``BING_SUBSCRIPTION_KEY``.
-
- .. code-block:: bash
-
- pip install -U langchain-community
- export BING_SUBSCRIPTION_KEY="your-api-key"
-
- Instantiation:
- .. code-block:: python
-
- from langchain_community.tools.bing_search import BingSearchResults
- from langchain_community.utilities import BingSearchAPIWrapper
-
- api_wrapper = BingSearchAPIWrapper()
- tool = BingSearchResults(api_wrapper=api_wrapper)
-
- Invocation with args:
- .. code-block:: python
-
- tool.invoke({"query": "what is the weather in SF?"})
-
- .. code-block:: python
-
- "[{'snippet': 'San Francisco, CAWeather Forecast, with current conditions, wind, air quality, and what to expect for the next 3 days.', 'title': 'San Francisco, CA Weather Forecast | AccuWeather', 'link': 'https://www.accuweather.com/en/us/san-francisco/94103/weather-forecast/347629'}, {'snippet': 'Tropical Storm Ernesto Forms; Fire Weather Concerns in the Great Basin: Hot Temperatures Return to the South-Central U.S. ... San Francisco CA 37.77°N 122.41°W (Elev. 131 ft) Last Update: 2:21 pm PDT Aug 12, 2024. Forecast Valid: 6pm PDT Aug 12, 2024-6pm PDT Aug 19, 2024 .', 'title': 'National Weather Service', 'link': 'https://forecast.weather.gov/zipcity.php?inputstring=San+Francisco,CA'}, {'snippet': 'Current weatherin San Francisco, CA. Check current conditions in San Francisco, CA with radar, hourly, and more.', 'title': 'San Francisco, CA Current Weather | AccuWeather', 'link': 'https://www.accuweather.com/en/us/san-francisco/94103/current-weather/347629'}, {'snippet': 'Everything you need to know about today's weatherin San Francisco, CA. High/Low, Precipitation Chances, Sunrise/Sunset, and today's Temperature History.', 'title': 'Weather Today for San Francisco, CA | AccuWeather', 'link': 'https://www.accuweather.com/en/us/san-francisco/94103/weather-today/347629'}]"
-
- Invocation with ToolCall:
-
- .. code-block:: python
-
- tool.invoke({"args": {"query":"what is the weather in SF?"}, "id": "1", "name": tool.name, "type": "tool_call"})
-
- .. code-block:: python
-
- ToolMessage(
- content="[{'snippet': 'Get the latest weather forecast for San Francisco, CA, including temperature, RealFeel, and chance of precipitation. Find out how the weather will affect your plans and activities in the city of ...', 'title': 'San Francisco, CA Weather Forecast | AccuWeather', 'link': 'https://www.accuweather.com/en/us/san-francisco/94103/weather-forecast/347629'}, {'snippet': 'Radar. Be prepared with the most accurate 10-day forecast for San Francisco, CA with highs, lows, chance of precipitation from The Weather Channel and Weather.com.', 'title': '10-Day Weather Forecast for San Francisco, CA - The Weather Channel', 'link': 'https://weather.com/weather/tenday/l/San+Francisco+CA+USCA0987:1:US'}, {'snippet': 'Tropical Storm Ernesto Forms; Fire Weather Concerns in the Great Basin: Hot Temperatures Return to the South-Central U.S. ... San Francisco CA 37.77°N 122.41°W (Elev. 131 ft) Last Update: 2:21 pm PDT Aug 12, 2024. Forecast Valid: 6pm PDT Aug 12, 2024-6pm PDT Aug 19, 2024 .', 'title': 'National Weather Service', 'link': 'https://forecast.weather.gov/zipcity.php?inputstring=San+Francisco,CA'}, {'snippet': 'Current weatherin San Francisco, CA. Check current conditions in San Francisco, CA with radar, hourly, and more.', 'title': 'San Francisco, CA Current Weather | AccuWeather', 'link': 'https://www.accuweather.com/en/us/san-francisco/94103/current-weather/347629'}]",
- artifact=[{'snippet': 'Get the latest weather forecast for San Francisco, CA, including temperature, RealFeel, and chance of precipitation. Find out how the weather will affect your plans and activities in the city of ...', 'title': 'San Francisco, CA Weather Forecast | AccuWeather', 'link': 'https://www.accuweather.com/en/us/san-francisco/94103/weather-forecast/347629'}, {'snippet': 'Radar. Be prepared with the most accurate 10-day forecast for San Francisco, CA with highs, lows, chance of precipitation from The Weather Channel and Weather.com.', 'title': '10-Day Weather Forecast for San Francisco, CA - The Weather Channel', 'link': 'https://weather.com/weather/tenday/l/San+Francisco+CA+USCA0987:1:US'}, {'snippet': 'Tropical Storm Ernesto Forms; Fire Weather Concerns in the Great Basin: Hot Temperatures Return to the South-Central U.S. ... San Francisco CA 37.77°N 122.41°W (Elev. 131 ft) Last Update: 2:21 pm PDT Aug 12, 2024. Forecast Valid: 6pm PDT Aug 12, 2024-6pm PDT Aug 19, 2024 .', 'title': 'National Weather Service', 'link': 'https://forecast.weather.gov/zipcity.php?inputstring=San+Francisco,CA'}, {'snippet': 'Current weatherin San Francisco, CA. Check current conditions in San Francisco, CA with radar, hourly, and more.', 'title': 'San Francisco, CA Current Weather | AccuWeather', 'link': 'https://www.accuweather.com/en/us/san-francisco/94103/current-weather/347629'}],
- name='bing_search_results_json',
- tool_call_id='1'
- )
-
- """ # noqa: E501
-
- name: str = "bing_search_results_json"
- description: str = (
- "A wrapper around Bing Search. "
- "Useful for when you need to answer questions about current events. "
- "Input should be a search query. Output is an array of the query results."
- )
- num_results: int = 4
- """Max search results to return, default is 4."""
- api_wrapper: BingSearchAPIWrapper
- response_format: Literal["content_and_artifact"] = "content_and_artifact"
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> Tuple[str, List[Dict]]:
- """Use the tool."""
- try:
- results = self.api_wrapper.results(query, self.num_results)
- return str(results), results
- except Exception as e:
- return repr(e), []
diff --git a/libs/community/langchain_community/tools/brave_search/__init__.py b/libs/community/langchain_community/tools/brave_search/__init__.py
deleted file mode 100644
index e69de29bb2..0000000000
diff --git a/libs/community/langchain_community/tools/brave_search/tool.py b/libs/community/langchain_community/tools/brave_search/tool.py
deleted file mode 100644
index c9d62210c2..0000000000
--- a/libs/community/langchain_community/tools/brave_search/tool.py
+++ /dev/null
@@ -1,92 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import Field, SecretStr
-
-from langchain_community.utilities.brave_search import BraveSearchWrapper
-
-
-class BraveSearch(BaseTool):
- """Tool that queries the BraveSearch.
-
- Api key can be provided as an environment variable BRAVE_SEARCH_API_KEY
- or as a parameter.
-
-
- Example usages:
- .. code-block:: python
- # uses BRAVE_SEARCH_API_KEY from environment
- tool = BraveSearch()
-
- .. code-block:: python
- # uses the provided api key
- tool = BraveSearch.from_api_key("your-api-key")
-
- .. code-block:: python
- # uses the provided api key and search kwargs
- tool = BraveSearch.from_api_key(
- api_key = "your-api-key",
- search_kwargs={"max_results": 5}
- )
-
- .. code-block:: python
- # uses BRAVE_SEARCH_API_KEY from environment
- tool = BraveSearch.from_search_kwargs({"max_results": 5})
- """
-
- name: str = "brave_search"
- description: str = (
- "a search engine. "
- "useful for when you need to answer questions about current events."
- " input should be a search query."
- )
- search_wrapper: BraveSearchWrapper = Field(default_factory=BraveSearchWrapper)
-
- @classmethod
- def from_api_key(
- cls, api_key: str, search_kwargs: Optional[dict] = None, **kwargs: Any
- ) -> BraveSearch:
- """Create a tool from an api key.
-
- Args:
- api_key: The api key to use.
- search_kwargs: Any additional kwargs to pass to the search wrapper.
- **kwargs: Any additional kwargs to pass to the tool.
-
- Returns:
- A tool.
- """
- wrapper = BraveSearchWrapper(
- api_key=SecretStr(api_key), search_kwargs=search_kwargs or {}
- )
- return cls(search_wrapper=wrapper, **kwargs)
-
- @classmethod
- def from_search_kwargs(cls, search_kwargs: dict, **kwargs: Any) -> BraveSearch:
- """Create a tool from search kwargs.
-
- Uses the environment variable BRAVE_SEARCH_API_KEY for api key.
-
- Args:
- search_kwargs: Any additional kwargs to pass to the search wrapper.
- **kwargs: Any additional kwargs to pass to the tool.
-
- Returns:
- A tool.
- """
- # we can not provide api key because it's calculated in the wrapper,
- # so the ignore is needed for linter
- # not ideal but needed to keep the tool code changes non-breaking
- wrapper = BraveSearchWrapper(search_kwargs=search_kwargs)
- return cls(search_wrapper=wrapper, **kwargs)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return self.search_wrapper.run(query)
diff --git a/libs/community/langchain_community/tools/cassandra_database/__init__.py b/libs/community/langchain_community/tools/cassandra_database/__init__.py
deleted file mode 100644
index 737e5be4ae..0000000000
--- a/libs/community/langchain_community/tools/cassandra_database/__init__.py
+++ /dev/null
@@ -1 +0,0 @@
-"""Cassandra Tool"""
diff --git a/libs/community/langchain_community/tools/cassandra_database/prompt.py b/libs/community/langchain_community/tools/cassandra_database/prompt.py
deleted file mode 100644
index ca264d3c59..0000000000
--- a/libs/community/langchain_community/tools/cassandra_database/prompt.py
+++ /dev/null
@@ -1,36 +0,0 @@
-"""Tools for interacting with an Apache Cassandra database."""
-
-QUERY_PATH_PROMPT = """"
-You are an Apache Cassandra expert query analysis bot with the following features
-and rules:
- - You will take a question from the end user about finding certain
- data in the database.
- - You will examine the schema of the database and create a query path.
- - You will provide the user with the correct query to find the data they are looking
- for showing the steps provided by the query path.
- - You will use best practices for querying Apache Cassandra using partition keys
- and clustering columns.
- - Avoid using ALLOW FILTERING in the query.
- - The goal is to find a query path, so it may take querying other tables to get
- to the final answer.
-
-The following is an example of a query path in JSON format:
-
- {
- "query_paths": [
- {
- "description": "Direct query to users table using email",
- "steps": [
- {
- "table": "user_credentials",
- "query":
- "SELECT userid FROM user_credentials WHERE email = 'example@example.com';"
- },
- {
- "table": "users",
- "query": "SELECT * FROM users WHERE userid = ?;"
- }
- ]
- }
- ]
-}"""
diff --git a/libs/community/langchain_community/tools/cassandra_database/tool.py b/libs/community/langchain_community/tools/cassandra_database/tool.py
deleted file mode 100644
index ab6e502fb0..0000000000
--- a/libs/community/langchain_community/tools/cassandra_database/tool.py
+++ /dev/null
@@ -1,142 +0,0 @@
-"""Tools for interacting with an Apache Cassandra database."""
-
-from __future__ import annotations
-
-import traceback
-from typing import TYPE_CHECKING, Any, Dict, Optional, Sequence, Type, Union
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, ConfigDict, Field
-
-from langchain_community.utilities.cassandra_database import CassandraDatabase
-
-if TYPE_CHECKING:
- from cassandra.cluster import ResultSet
-
-
-class BaseCassandraDatabaseTool(BaseModel):
- """Base tool for interacting with an Apache Cassandra database."""
-
- db: CassandraDatabase = Field(exclude=True)
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- )
-
-
-class _QueryCassandraDatabaseToolInput(BaseModel):
- query: str = Field(..., description="A detailed and correct CQL query.")
-
-
-class QueryCassandraDatabaseTool(BaseCassandraDatabaseTool, BaseTool):
- """Tool for querying an Apache Cassandra database with provided CQL."""
-
- name: str = "cassandra_db_query"
- description: str = """
- Execute a CQL query against the database and get back the result.
- If the query is not correct, an error message will be returned.
- If an error is returned, rewrite the query, check the query, and try again.
- """
- args_schema: Type[BaseModel] = _QueryCassandraDatabaseToolInput
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> Union[str, Sequence[Dict[str, Any]], ResultSet]:
- """Execute the query, return the results or an error message."""
- try:
- return self.db.run(query)
- except Exception as e:
- """Format the error message"""
- return f"Error: {e}\n{traceback.format_exc()}"
-
-
-class _GetSchemaCassandraDatabaseToolInput(BaseModel):
- keyspace: str = Field(
- ...,
- description=("The name of the keyspace for which to return the schema."),
- )
-
-
-class GetSchemaCassandraDatabaseTool(BaseCassandraDatabaseTool, BaseTool):
- """Tool for getting the schema of a keyspace in an Apache Cassandra database."""
-
- name: str = "cassandra_db_schema"
- description: str = """
- Input to this tool is a keyspace name, output is a table description
- of Apache Cassandra tables.
- If the query is not correct, an error message will be returned.
- If an error is returned, report back to the user that the keyspace
- doesn't exist and stop.
- """
-
- args_schema: Type[BaseModel] = _GetSchemaCassandraDatabaseToolInput
-
- def _run(
- self,
- keyspace: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Get the schema for a keyspace."""
- try:
- tables = self.db.get_keyspace_tables(keyspace)
- return "".join([table.as_markdown() + "\n\n" for table in tables])
- except Exception as e:
- """Format the error message"""
- return f"Error: {e}\n{traceback.format_exc()}"
-
-
-class _GetTableDataCassandraDatabaseToolInput(BaseModel):
- keyspace: str = Field(
- ...,
- description=("The name of the keyspace containing the table."),
- )
- table: str = Field(
- ...,
- description=("The name of the table for which to return data."),
- )
- predicate: str = Field(
- ...,
- description=("The predicate for the query that uses the primary key."),
- )
- limit: int = Field(
- ...,
- description=("The maximum number of rows to return."),
- )
-
-
-class GetTableDataCassandraDatabaseTool(BaseCassandraDatabaseTool, BaseTool):
- """
- Tool for getting data from a table in an Apache Cassandra database.
- Use the WHERE clause to specify the predicate for the query that uses the
- primary key. A blank predicate will return all rows. Avoid this if possible.
- Use the limit to specify the number of rows to return. A blank limit will
- return all rows.
- """
-
- name: str = "cassandra_db_select_table_data"
- description: str = """
- Tool for getting data from a table in an Apache Cassandra database.
- Use the WHERE clause to specify the predicate for the query that uses the
- primary key. A blank predicate will return all rows. Avoid this if possible.
- Use the limit to specify the number of rows to return. A blank limit will
- return all rows.
- """
- args_schema: Type[BaseModel] = _GetTableDataCassandraDatabaseToolInput
-
- def _run(
- self,
- keyspace: str,
- table: str,
- predicate: str,
- limit: int,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Get data from a table in a keyspace."""
- try:
- return self.db.get_table_data(keyspace, table, predicate, limit)
- except Exception as e:
- """Format the error message"""
- return f"Error: {e}\n{traceback.format_exc()}"
diff --git a/libs/community/langchain_community/tools/clickup/__init__.py b/libs/community/langchain_community/tools/clickup/__init__.py
deleted file mode 100644
index e69de29bb2..0000000000
diff --git a/libs/community/langchain_community/tools/clickup/prompt.py b/libs/community/langchain_community/tools/clickup/prompt.py
deleted file mode 100644
index 5f1f51dbf7..0000000000
--- a/libs/community/langchain_community/tools/clickup/prompt.py
+++ /dev/null
@@ -1,131 +0,0 @@
-# flake8: noqa
-CLICKUP_TASK_CREATE_PROMPT = """
- This tool is a wrapper around clickup's create_task API, useful when you need to create a CLICKUP task.
- The input to this tool is a dictionary specifying the fields of the CLICKUP task, and will be passed into clickup's CLICKUP `create_task` function.
- Only add fields described by the user.
- Use the following mapping in order to map the user's priority to the clickup priority: {{
- Urgent = 1,
- High = 2,
- Normal = 3,
- Low = 4,
- }}. If the user passes in "urgent" replace the priority value as 1.
-
- Here are a few task descriptions and corresponding input examples:
- Task: create a task called "Daily report"
- Example Input: {{"name": "Daily report"}}
- Task: Make an open task called "ClickUp toolkit refactor" with description "Refactor the clickup toolkit to use dataclasses for parsing", with status "open"
- Example Input: {{"name": "ClickUp toolkit refactor", "description": "Refactor the clickup toolkit to use dataclasses for parsing", "status": "Open"}}
- Task: create a task with priority 3 called "New Task Name" with description "New Task Description", with status "open"
- Example Input: {{"name": "New Task Name", "description": "New Task Description", "status": "Open", "priority": 3}}
- Task: Add a task called "Bob's task" and assign it to Bob (user id: 81928627)
- Example Input: {{"name": "Bob's task", "description": "Task for Bob", "assignees": [81928627]}}
- """
-
-CLICKUP_LIST_CREATE_PROMPT = """
- This tool is a wrapper around clickup's create_list API, useful when you need to create a CLICKUP list.
- The input to this tool is a dictionary specifying the fields of a clickup list, and will be passed to clickup's create_list function.
- Only add fields described by the user.
- Use the following mapping in order to map the user's priority to the clickup priority: {{
- Urgent = 1,
- High = 2,
- Normal = 3,
- Low = 4,
- }}. If the user passes in "urgent" replace the priority value as 1.
-
- Here are a few list descriptions and corresponding input examples:
- Description: make a list with name "General List"
- Example Input: {{"name": "General List"}}
- Description: add a new list ("TODOs") with low priority
- Example Input: {{"name": "General List", "priority": 4}}
- Description: create a list with name "List name", content "List content", priority 2, and status "red"
- Example Input: {{"name": "List name", "content": "List content", "priority": 2, "status": "red"}}
-"""
-
-CLICKUP_FOLDER_CREATE_PROMPT = """
- This tool is a wrapper around clickup's create_folder API, useful when you need to create a CLICKUP folder.
- The input to this tool is a dictionary specifying the fields of a clickup folder, and will be passed to clickup's create_folder function.
- For example, to create a folder with name "Folder name" you would pass in the following dictionary:
- {{
- "name": "Folder name",
- }}
-"""
-
-CLICKUP_GET_TASK_PROMPT = """
- This tool is a wrapper around clickup's API,
- Do NOT use to get a task specific attribute. Use get task attribute instead.
- useful when you need to get a specific task for the user. Given the task id you want to create a request similar to the following dictionary:
- payload = {{"task_id": "86a0t44tq"}}
- """
-
-CLICKUP_GET_TASK_ATTRIBUTE_PROMPT = """
- This tool is a wrapper around clickup's API,
- useful when you need to get a specific attribute from a task. Given the task id and desired attribute create a request similar to the following dictionary:
- payload = {{"task_id": "", "attribute_name": ""}}
-
- Here are some example queries their corresponding payloads:
- Get the name of task 23jn23kjn -> {{"task_id": "23jn23kjn", "attribute_name": "name"}}
- What is the priority of task 86a0t44tq? -> {{"task_id": "86a0t44tq", "attribute_name": "priority"}}
- Output the description of task sdc9ds9jc -> {{"task_id": "sdc9ds9jc", "attribute_name": "description"}}
- Who is assigned to task bgjfnbfg0 -> {{"task_id": "bgjfnbfg0", "attribute_name": "assignee"}}
- Which is the status of task kjnsdcjc? -> {{"task_id": "kjnsdcjc", "attribute_name": "description"}}
- How long is the time estimate of task sjncsd999? -> {{"task_id": "sjncsd999", "attribute_name": "time_estimate"}}
- Is task jnsd98sd archived?-> {{"task_id": "jnsd98sd", "attribute_name": "archive"}}
- """
-
-CLICKUP_GET_ALL_TEAMS_PROMPT = """
- This tool is a wrapper around clickup's API, useful when you need to get all teams that the user is a part of.
- To get a list of all the teams there is no necessary request parameters.
- """
-
-CLICKUP_GET_LIST_PROMPT = """
- This tool is a wrapper around clickup's API,
- useful when you need to get a specific list for the user. Given the list id you want to create a request similar to the following dictionary:
- payload = {{"list_id": "901300608424"}}
- """
-
-CLICKUP_GET_FOLDERS_PROMPT = """
- This tool is a wrapper around clickup's API,
- useful when you need to get a specific folder for the user. Given the user's workspace id you want to create a request similar to the following dictionary:
- payload = {{"folder_id": "90130119692"}}
- """
-
-CLICKUP_GET_SPACES_PROMPT = """
- This tool is a wrapper around clickup's API,
- useful when you need to get all the spaces available to a user. Given the user's workspace id you want to create a request similar to the following dictionary:
- payload = {{"team_id": "90130119692"}}
- """
-
-CLICKUP_GET_SPACES_PROMPT = """
- This tool is a wrapper around clickup's API,
- useful when you need to get all the spaces available to a user. Given the user's workspace id you want to create a request similar to the following dictionary:
- payload = {{"team_id": "90130119692"}}
- """
-
-CLICKUP_UPDATE_TASK_PROMPT = """
- This tool is a wrapper around clickup's API,
- useful when you need to update a specific attribute of a task. Given the task id, desired attribute to change and the new value you want to create a request similar to the following dictionary:
- payload = {{"task_id": "", "attribute_name": "", "value": ""}}
-
- Here are some example queries their corresponding payloads:
- Change the name of task 23jn23kjn to new task name -> {{"task_id": "23jn23kjn", "attribute_name": "name", "value": "new task name"}}
- Update the priority of task 86a0t44tq to 1 -> {{"task_id": "86a0t44tq", "attribute_name": "priority", "value": 1}}
- Re-write the description of task sdc9ds9jc to 'a new task description' -> {{"task_id": "sdc9ds9jc", "attribute_name": "description", "value": "a new task description"}}
- Forward the status of task kjnsdcjc to done -> {{"task_id": "kjnsdcjc", "attribute_name": "description", "status": "done"}}
- Increase the time estimate of task sjncsd999 to 3h -> {{"task_id": "sjncsd999", "attribute_name": "time_estimate", "value": 8000}}
- Archive task jnsd98sd -> {{"task_id": "jnsd98sd", "attribute_name": "archive", "value": true}}
- *IMPORTANT*: Pay attention to the exact syntax above and the correct use of quotes.
- For changing priority and time estimates, we expect integers (int).
- For name, description and status we expect strings (str).
- For archive, we expect a boolean (bool).
- """
-
-CLICKUP_UPDATE_TASK_ASSIGNEE_PROMPT = """
- This tool is a wrapper around clickup's API,
- useful when you need to update the assignees of a task. Given the task id, the operation add or remove (rem), and the list of user ids. You want to create a request similar to the following dictionary:
- payload = {{"task_id": "", "operation": "", "users": [, ]}}
-
- Here are some example queries their corresponding payloads:
- Add 81928627 and 3987234 as assignees to task 21hw21jn -> {{"task_id": "21hw21jn", "operation": "add", "users": [81928627, 3987234]}}
- Remove 67823487 as assignee from task jin34ji4 -> {{"task_id": "jin34ji4", "operation": "rem", "users": [67823487]}}
- *IMPORTANT*: Users id should always be ints.
- """
diff --git a/libs/community/langchain_community/tools/clickup/tool.py b/libs/community/langchain_community/tools/clickup/tool.py
deleted file mode 100644
index 03b3fde586..0000000000
--- a/libs/community/langchain_community/tools/clickup/tool.py
+++ /dev/null
@@ -1,43 +0,0 @@
-"""
-This tool allows agents to interact with the clickup library
-and operate on a Clickup instance.
-To use this tool, you must first set as environment variables:
- client_secret
- client_id
- code
-
-Below is a sample script that uses the Clickup tool:
-
-```python
-from langchain_community.agent_toolkits.clickup.toolkit import ClickupToolkit
-from langchain_community.utilities.clickup import ClickupAPIWrapper
-
-clickup = ClickupAPIWrapper()
-toolkit = ClickupToolkit.from_clickup_api_wrapper(clickup)
-```
-"""
-
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import Field
-
-from langchain_community.utilities.clickup import ClickupAPIWrapper
-
-
-class ClickupAction(BaseTool):
- """Tool that queries the Clickup API."""
-
- api_wrapper: ClickupAPIWrapper = Field(default_factory=ClickupAPIWrapper)
- mode: str
- name: str = ""
- description: str = ""
-
- def _run(
- self,
- instructions: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Clickup API to run an operation."""
- return self.api_wrapper.run(self.mode, instructions)
diff --git a/libs/community/langchain_community/tools/cogniswitch/__init__.py b/libs/community/langchain_community/tools/cogniswitch/__init__.py
deleted file mode 100644
index 3a89a8d7d3..0000000000
--- a/libs/community/langchain_community/tools/cogniswitch/__init__.py
+++ /dev/null
@@ -1 +0,0 @@
-"Cogniswitch Tools"
diff --git a/libs/community/langchain_community/tools/cogniswitch/tool.py b/libs/community/langchain_community/tools/cogniswitch/tool.py
deleted file mode 100644
index 41c9663fb5..0000000000
--- a/libs/community/langchain_community/tools/cogniswitch/tool.py
+++ /dev/null
@@ -1,399 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Dict, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-
-
-class CogniswitchKnowledgeRequest(BaseTool):
- """Tool that uses the Cogniswitch service to answer questions.
-
- name: str = "cogniswitch_knowledge_request"
- description: str = (
- "A wrapper around cogniswitch service to answer the question
- from the knowledge base."
- "Input should be a search query."
- )
- """
-
- name: str = "cogniswitch_knowledge_request"
- description: str = """A wrapper around cogniswitch service to
- answer the question from the knowledge base."""
- cs_token: str
- OAI_token: str
- apiKey: str
- api_url: str = "https://api.cogniswitch.ai:8243/cs-api/0.0.1/cs/knowledgeRequest"
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> Dict[str, Any]:
- """
- Use the tool to answer a query.
-
- Args:
- query (str): Natural language query,
- that you would like to ask to your knowledge graph.
- run_manager (Optional[CallbackManagerForChainRun]):
- Manager for chain run callbacks.
-
- Returns:
- Dict[str, Any]: Output dictionary containing
- the 'response' from the service.
- """
- response = self.answer_cs(self.cs_token, self.OAI_token, query, self.apiKey)
- return response
-
- def answer_cs(self, cs_token: str, OAI_token: str, query: str, apiKey: str) -> dict:
- """
- Send a query to the Cogniswitch service and retrieve the response.
-
- Args:
- cs_token (str): Cogniswitch token.
- OAI_token (str): OpenAI token.
- apiKey (str): OAuth token.
- query (str): Query to be answered.
-
- Returns:
- dict: Response JSON from the Cogniswitch service.
- """
- if not cs_token:
- raise ValueError("Missing cs_token")
- if not OAI_token:
- raise ValueError("Missing OpenAI token")
- if not apiKey:
- raise ValueError("Missing cogniswitch OAuth token")
- if not query:
- raise ValueError("Missing input query")
-
- headers = {
- "apiKey": apiKey,
- "platformToken": cs_token,
- "openAIToken": OAI_token,
- }
-
- data = {"query": query}
- response = requests.post(self.api_url, headers=headers, verify=False, data=data)
- return response.json()
-
-
-class CogniswitchKnowledgeStatus(BaseTool):
- """Tool that uses the Cogniswitch services to get the
- status of the document or url uploaded.
-
- name: str = "cogniswitch_knowledge_status"
- description: str = (
- "A wrapper around cogniswitch services to know the status of
- the document uploaded from a url or a file. "
- "Input should be a file name or the url link"
- )
- """
-
- name: str = "cogniswitch_knowledge_status"
- description: str = """A wrapper around cogniswitch services to know
- the status of the document uploaded from a url or a file."""
- cs_token: str
- OAI_token: str
- apiKey: str
- knowledge_status_url: str = (
- "https://api.cogniswitch.ai:8243/cs-api/0.0.1/cs/knowledgeSource/status"
- )
-
- def _run(
- self,
- document_name: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> Dict[str, Any]:
- """
- Use the tool to know the status of the document uploaded.
-
- Args:
- document_name (str): name of the document or
- the url uploaded
- run_manager (Optional[CallbackManagerForChainRun]):
- Manager for chain run callbacks.
-
- Returns:
- Dict[str, Any]: Output dictionary containing
- the 'response' from the service.
- """
- response = self.knowledge_status(document_name)
- return response
-
- def knowledge_status(self, document_name: str) -> dict:
- """
- Use this function to know the status of the document or the URL uploaded
- Args:
- document_name (str): The document name or the url that is uploaded.
-
- Returns:
- dict: Response JSON from the Cogniswitch service.
- """
-
- params = {"docName": document_name, "platformToken": self.cs_token}
- headers = {
- "apiKey": self.apiKey,
- "openAIToken": self.OAI_token,
- "platformToken": self.cs_token,
- }
- response = requests.get(
- self.knowledge_status_url,
- headers=headers,
- params=params,
- verify=False,
- )
- if response.status_code == 200:
- source_info = response.json()
- source_data = dict(source_info[-1])
- status = source_data.get("status")
- if status == 0:
- source_data["status"] = "SUCCESS"
- elif status == 1:
- source_data["status"] = "PROCESSING"
- elif status == 2:
- source_data["status"] = "UPLOADED"
- elif status == 3:
- source_data["status"] = "FAILURE"
- elif status == 4:
- source_data["status"] = "UPLOAD_FAILURE"
- elif status == 5:
- source_data["status"] = "REJECTED"
-
- if "filePath" in source_data.keys():
- source_data.pop("filePath")
- if "savedFileName" in source_data.keys():
- source_data.pop("savedFileName")
- if "integrationConfigId" in source_data.keys():
- source_data.pop("integrationConfigId")
- if "metaData" in source_data.keys():
- source_data.pop("metaData")
- if "docEntryId" in source_data.keys():
- source_data.pop("docEntryId")
- return source_data
- else:
- # error_message = response.json()["message"]
- return {
- "message": response.status_code,
- }
-
-
-class CogniswitchKnowledgeSourceFile(BaseTool):
- """Tool that uses the Cogniswitch services to store data from file.
-
- name: str = "cogniswitch_knowledge_source_file"
- description: str = (
- "This calls the CogniSwitch services to analyze & store data from a file.
- If the input looks like a file path, assign that string value to file key.
- Assign document name & description only if provided in input."
- )
- """
-
- name: str = "cogniswitch_knowledge_source_file"
- description: str = """
- This calls the CogniSwitch services to analyze & store data from a file.
- If the input looks like a file path, assign that string value to file key.
- Assign document name & description only if provided in input.
- """
- cs_token: str
- OAI_token: str
- apiKey: str
- knowledgesource_file: str = (
- "https://api.cogniswitch.ai:8243/cs-api/0.0.1/cs/knowledgeSource/file"
- )
-
- def _run(
- self,
- file: Optional[str] = None,
- document_name: Optional[str] = None,
- document_description: Optional[str] = None,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> Dict[str, Any]:
- """
- Execute the tool to store the data given from a file.
- This calls the CogniSwitch services to analyze & store data from a file.
- If the input looks like a file path, assign that string value to file key.
- Assign document name & description only if provided in input.
-
- Args:
- file Optional[str]: The file path of your knowledge
- document_name Optional[str]: Name of your knowledge document
- document_description Optional[str]: Description of your knowledge document
- run_manager (Optional[CallbackManagerForChainRun]):
- Manager for chain run callbacks.
-
- Returns:
- Dict[str, Any]: Output dictionary containing
- the 'response' from the service.
- """
- if not file:
- return {
- "message": "No input provided",
- }
- else:
- response = self.store_data(
- file=file,
- document_name=document_name,
- document_description=document_description,
- )
- return response
-
- def store_data(
- self,
- file: Optional[str],
- document_name: Optional[str],
- document_description: Optional[str],
- ) -> dict:
- """
- Store data using the Cogniswitch service.
- This calls the CogniSwitch services to analyze & store data from a file.
- If the input looks like a file path, assign that string value to file key.
- Assign document name & description only if provided in input.
-
- Args:
- file (Optional[str]): file path of your file.
- the current files supported by the files are
- .txt, .pdf, .docx, .doc, .html
- document_name (Optional[str]): Name of the document you are uploading.
- document_description (Optional[str]): Description of the document.
-
- Returns:
- dict: Response JSON from the Cogniswitch service.
- """
- headers = {
- "apiKey": self.apiKey,
- "openAIToken": self.OAI_token,
- "platformToken": self.cs_token,
- }
- data: Dict[str, Any]
- if not document_name:
- document_name = ""
- if not document_description:
- document_description = ""
-
- if file is not None:
- files = {"file": open(file, "rb")}
-
- data = {
- "documentName": document_name,
- "documentDescription": document_description,
- }
- response = requests.post(
- self.knowledgesource_file,
- headers=headers,
- verify=False,
- data=data,
- files=files,
- )
- if response.status_code == 200:
- return response.json()
- else:
- return {"message": "Bad Request"}
-
-
-class CogniswitchKnowledgeSourceURL(BaseTool):
- """Tool that uses the Cogniswitch services to store data from a URL.
-
- name: str = "cogniswitch_knowledge_source_url"
- description: str = (
- "This calls the CogniSwitch services to analyze & store data from a url.
- the URL is provided in input, assign that value to the url key.
- Assign document name & description only if provided in input"
- )
- """
-
- name: str = "cogniswitch_knowledge_source_url"
- description: str = """
- This calls the CogniSwitch services to analyze & store data from a url.
- the URL is provided in input, assign that value to the url key.
- Assign document name & description only if provided in input"""
- cs_token: str
- OAI_token: str
- apiKey: str
- knowledgesource_url: str = (
- "https://api.cogniswitch.ai:8243/cs-api/0.0.1/cs/knowledgeSource/url"
- )
-
- def _run(
- self,
- url: Optional[str] = None,
- document_name: Optional[str] = None,
- document_description: Optional[str] = None,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> Dict[str, Any]:
- """
- Execute the tool to store the data given from a url.
- This calls the CogniSwitch services to analyze & store data from a url.
- the URL is provided in input, assign that value to the url key.
- Assign document name & description only if provided in input.
-
- Args:
- url Optional[str]: The website/url link of your knowledge
- document_name Optional[str]: Name of your knowledge document
- document_description Optional[str]: Description of your knowledge document
- run_manager (Optional[CallbackManagerForChainRun]):
- Manager for chain run callbacks.
-
- Returns:
- Dict[str, Any]: Output dictionary containing
- the 'response' from the service.
- """
- if not url:
- return {
- "message": "No input provided",
- }
- response = self.store_data(
- url=url,
- document_name=document_name,
- document_description=document_description,
- )
- return response
-
- def store_data(
- self,
- url: Optional[str],
- document_name: Optional[str],
- document_description: Optional[str],
- ) -> dict:
- """
- Store data using the Cogniswitch service.
- This calls the CogniSwitch services to analyze & store data from a url.
- the URL is provided in input, assign that value to the url key.
- Assign document name & description only if provided in input.
-
- Args:
- url (Optional[str]): URL link.
- document_name (Optional[str]): Name of the document you are uploading.
- document_description (Optional[str]): Description of the document.
-
- Returns:
- dict: Response JSON from the Cogniswitch service.
- """
- headers = {
- "apiKey": self.apiKey,
- "openAIToken": self.OAI_token,
- "platformToken": self.cs_token,
- }
- data: Dict[str, Any]
- if not document_name:
- document_name = ""
- if not document_description:
- document_description = ""
- if not url:
- return {
- "message": "No input provided",
- }
- else:
- data = {"url": url}
- response = requests.post(
- self.knowledgesource_url,
- headers=headers,
- verify=False,
- data=data,
- )
- if response.status_code == 200:
- return response.json()
- else:
- return {"message": "Bad Request"}
diff --git a/libs/community/langchain_community/tools/connery/__init__.py b/libs/community/langchain_community/tools/connery/__init__.py
deleted file mode 100644
index 1fcf2760ba..0000000000
--- a/libs/community/langchain_community/tools/connery/__init__.py
+++ /dev/null
@@ -1,8 +0,0 @@
-"""
-This module contains the ConneryAction Tool and ConneryService.
-"""
-
-from .service import ConneryService
-from .tool import ConneryAction
-
-__all__ = ["ConneryAction", "ConneryService"]
diff --git a/libs/community/langchain_community/tools/connery/models.py b/libs/community/langchain_community/tools/connery/models.py
deleted file mode 100644
index 537f58bea1..0000000000
--- a/libs/community/langchain_community/tools/connery/models.py
+++ /dev/null
@@ -1,32 +0,0 @@
-from typing import Any, List, Optional
-
-from pydantic import BaseModel
-
-
-class Validation(BaseModel):
- """Connery Action parameter validation model."""
-
- required: Optional[bool] = None
-
-
-class Parameter(BaseModel):
- """Connery Action parameter model."""
-
- key: str
- title: str
- description: Optional[str] = None
- type: Any
- validation: Optional[Validation] = None
-
-
-class Action(BaseModel):
- """Connery Action model."""
-
- id: str
- key: str
- title: str
- description: Optional[str] = None
- type: str
- inputParameters: List[Parameter]
- outputParameters: List[Parameter]
- pluginId: str
diff --git a/libs/community/langchain_community/tools/connery/service.py b/libs/community/langchain_community/tools/connery/service.py
deleted file mode 100644
index bbe4cc8c18..0000000000
--- a/libs/community/langchain_community/tools/connery/service.py
+++ /dev/null
@@ -1,166 +0,0 @@
-import json
-from typing import Any, Dict, List, Optional
-
-import requests
-from langchain_core.utils.env import get_from_dict_or_env
-from pydantic import BaseModel, model_validator
-
-from langchain_community.tools.connery.models import Action
-from langchain_community.tools.connery.tool import ConneryAction
-
-
-class ConneryService(BaseModel):
- """Service for interacting with the Connery Runner API.
-
- It gets the list of available actions from the Connery Runner,
- wraps them in ConneryAction Tools and returns them to the user.
- It also provides a method for running the actions.
- """
-
- runner_url: Optional[str] = None
- api_key: Optional[str] = None
-
- @model_validator(mode="before")
- @classmethod
- def validate_attributes(cls, values: Dict) -> Any:
- """
- Validate the attributes of the ConneryService class.
- Parameters:
- values (dict): The arguments to validate.
- Returns:
- dict: The validated arguments.
- """
-
- runner_url = get_from_dict_or_env(values, "runner_url", "CONNERY_RUNNER_URL")
- api_key = get_from_dict_or_env(values, "api_key", "CONNERY_RUNNER_API_KEY")
-
- if not runner_url:
- raise ValueError("CONNERY_RUNNER_URL environment variable must be set.")
- if not api_key:
- raise ValueError("CONNERY_RUNNER_API_KEY environment variable must be set.")
-
- values["runner_url"] = runner_url
- values["api_key"] = api_key
-
- return values
-
- def list_actions(self) -> List[ConneryAction]:
- """
- Returns the list of actions available in the Connery Runner.
- Returns:
- List[ConneryAction]: The list of actions available in the Connery Runner.
- """
-
- return [
- ConneryAction.create_instance(action, self)
- for action in self._list_actions()
- ]
-
- def get_action(self, action_id: str) -> ConneryAction:
- """
- Returns the specified action available in the Connery Runner.
- Parameters:
- action_id (str): The ID of the action to return.
- Returns:
- ConneryAction: The action with the specified ID.
- """
-
- return ConneryAction.create_instance(self._get_action(action_id), self)
-
- def run_action(self, action_id: str, input: Dict[str, str] = {}) -> Dict[str, str]:
- """
- Runs the specified Connery Action with the provided input.
- Parameters:
- action_id (str): The ID of the action to run.
- input (Dict[str, str]): The input object expected by the action.
- Returns:
- Dict[str, str]: The output of the action.
- """
-
- return self._run_action(action_id, input)
-
- def _list_actions(self) -> List[Action]:
- """
- Returns the list of actions available in the Connery Runner.
- Returns:
- List[Action]: The list of actions available in the Connery Runner.
- """
-
- response = requests.get(
- f"{self.runner_url}/v1/actions", headers=self._get_headers()
- )
-
- if not response.ok:
- raise ValueError(
- (
- "Failed to list actions."
- f"Status code: {response.status_code}."
- f"Error message: {response.json()['error']['message']}"
- )
- )
-
- return [Action(**action) for action in response.json()["data"]]
-
- def _get_action(self, action_id: str) -> Action:
- """
- Returns the specified action available in the Connery Runner.
- Parameters:
- action_id (str): The ID of the action to return.
- Returns:
- Action: The action with the specified ID.
- """
-
- actions = self._list_actions()
- action = next((action for action in actions if action.id == action_id), None)
- if not action:
- raise ValueError(
- (
- f"The action with ID {action_id} was not found in the list"
- "of available actions in the Connery Runner."
- )
- )
- return action
-
- def _run_action(self, action_id: str, input: Dict[str, str] = {}) -> Dict[str, str]:
- """
- Runs the specified Connery Action with the provided input.
- Parameters:
- action_id (str): The ID of the action to run.
- prompt (str): This is a plain English prompt
- with all the information needed to run the action.
- input (Dict[str, str]): The input object expected by the action.
- If provided together with the prompt,
- the input takes precedence over the input specified in the prompt.
- Returns:
- Dict[str, str]: The output of the action.
- """
-
- response = requests.post(
- f"{self.runner_url}/v1/actions/{action_id}/run",
- headers=self._get_headers(),
- data=json.dumps({"input": input}),
- )
-
- if not response.ok:
- raise ValueError(
- (
- "Failed to run action."
- f"Status code: {response.status_code}."
- f"Error message: {response.json()['error']['message']}"
- )
- )
-
- if not response.json()["data"]["output"]:
- return {}
- else:
- return response.json()["data"]["output"]
-
- def _get_headers(self) -> Dict[str, str]:
- """
- Returns a standard set of HTTP headers
- to be used in API calls to the Connery runner.
- Returns:
- Dict[str, str]: The standard set of HTTP headers.
- """
-
- return {"Content-Type": "application/json", "x-api-key": self.api_key or ""}
diff --git a/libs/community/langchain_community/tools/connery/tool.py b/libs/community/langchain_community/tools/connery/tool.py
deleted file mode 100644
index d74bfc1743..0000000000
--- a/libs/community/langchain_community/tools/connery/tool.py
+++ /dev/null
@@ -1,162 +0,0 @@
-import asyncio
-from functools import partial
-from typing import Any, Dict, List, Optional, Type
-
-from langchain_core.callbacks.manager import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field, create_model, model_validator
-
-from langchain_community.tools.connery.models import Action, Parameter
-
-
-class ConneryAction(BaseTool):
- """Connery Action tool."""
-
- name: str
- description: str
- args_schema: Type[BaseModel]
-
- action: Action
- connery_service: Any
-
- def _run(
- self,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- **kwargs: Any,
- ) -> Dict[str, str]:
- """
- Runs the Connery Action with the provided input.
- Parameters:
- kwargs (Dict[str, str]): The input dictionary expected by the action.
- Returns:
- Dict[str, str]: The output of the action.
- """
-
- return self.connery_service.run_action(self.action.id, kwargs)
-
- async def _arun(
- self,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- **kwargs: Any,
- ) -> Dict[str, str]:
- """
- Runs the Connery Action asynchronously with the provided input.
- Parameters:
- kwargs (Dict[str, str]): The input dictionary expected by the action.
- Returns:
- Dict[str, str]: The output of the action.
- """
-
- func = partial(self._run, **kwargs)
- return await asyncio.get_event_loop().run_in_executor(None, func)
-
- def get_schema_json(self) -> str:
- """
- Returns the JSON representation of the Connery Action Tool schema.
- This is useful for debugging.
- Returns:
- str: The JSON representation of the Connery Action Tool schema.
- """
-
- return self.args_schema.schema_json(indent=2)
-
- @model_validator(mode="before")
- @classmethod
- def validate_attributes(cls, values: dict) -> Any:
- """
- Validate the attributes of the ConneryAction class.
- Parameters:
- values (dict): The arguments to validate.
- Returns:
- dict: The validated arguments.
- """
-
- # Import ConneryService here and check if it is an instance
- # of ConneryService to avoid circular imports
- from .service import ConneryService
-
- if not isinstance(values.get("connery_service"), ConneryService):
- raise ValueError(
- "The attribute 'connery_service' must be an instance of ConneryService."
- )
-
- if not values.get("name"):
- raise ValueError("The attribute 'name' must be set.")
- if not values.get("description"):
- raise ValueError("The attribute 'description' must be set.")
- if not values.get("args_schema"):
- raise ValueError("The attribute 'args_schema' must be set.")
- if not values.get("action"):
- raise ValueError("The attribute 'action' must be set.")
- if not values.get("connery_service"):
- raise ValueError("The attribute 'connery_service' must be set.")
-
- return values
-
- @classmethod
- def create_instance(cls, action: Action, connery_service: Any) -> "ConneryAction":
- """
- Creates a Connery Action Tool from a Connery Action.
- Parameters:
- action (Action): The Connery Action to wrap in a Connery Action Tool.
- connery_service (ConneryService): The Connery Service
- to run the Connery Action. We use Any here to avoid circular imports.
- Returns:
- ConneryAction: The Connery Action Tool.
- """
-
- # Import ConneryService here and check if it is an instance
- # of ConneryService to avoid circular imports
- from .service import ConneryService
-
- if not isinstance(connery_service, ConneryService):
- raise ValueError(
- "The connery_service must be an instance of ConneryService."
- )
-
- input_schema = cls._create_input_schema(action.inputParameters)
- description = action.title + (
- ": " + action.description if action.description else ""
- )
-
- instance = cls(
- name=action.id,
- description=description,
- args_schema=input_schema,
- action=action,
- connery_service=connery_service,
- )
-
- return instance
-
- @classmethod
- def _create_input_schema(cls, inputParameters: List[Parameter]) -> Type[BaseModel]:
- """
- Creates an input schema for a Connery Action Tool
- based on the input parameters of the Connery Action.
- Parameters:
- inputParameters: List of input parameters of the Connery Action.
- Returns:
- Type[BaseModel]: The input schema for the Connery Action Tool.
- """
-
- dynamic_input_fields: Dict[str, Any] = {}
-
- for param in inputParameters:
- default = ... if param.validation and param.validation.required else None
- title = param.title
- description = param.title + (
- ": " + param.description if param.description else ""
- )
- type = param.type
-
- dynamic_input_fields[param.key] = (
- type,
- Field(default, title=title, description=description),
- )
-
- InputModel = create_model("InputSchema", **dynamic_input_fields)
- return InputModel
diff --git a/libs/community/langchain_community/tools/convert_to_openai.py b/libs/community/langchain_community/tools/convert_to_openai.py
deleted file mode 100644
index 423f927127..0000000000
--- a/libs/community/langchain_community/tools/convert_to_openai.py
+++ /dev/null
@@ -1,6 +0,0 @@
-from langchain_core.utils.function_calling import (
- format_tool_to_openai_function,
- format_tool_to_openai_tool,
-)
-
-__all__ = ["format_tool_to_openai_function", "format_tool_to_openai_tool"]
diff --git a/libs/community/langchain_community/tools/databricks/__init__.py b/libs/community/langchain_community/tools/databricks/__init__.py
deleted file mode 100644
index 9a1d5ffe53..0000000000
--- a/libs/community/langchain_community/tools/databricks/__init__.py
+++ /dev/null
@@ -1,3 +0,0 @@
-from langchain_community.tools.databricks.tool import UCFunctionToolkit
-
-__all__ = ["UCFunctionToolkit"]
diff --git a/libs/community/langchain_community/tools/databricks/_execution.py b/libs/community/langchain_community/tools/databricks/_execution.py
deleted file mode 100644
index 67ab2e7c26..0000000000
--- a/libs/community/langchain_community/tools/databricks/_execution.py
+++ /dev/null
@@ -1,254 +0,0 @@
-import inspect
-import json
-import logging
-import os
-import time
-from dataclasses import dataclass
-from io import StringIO
-from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional
-
-if TYPE_CHECKING:
- from databricks.sdk import WorkspaceClient
- from databricks.sdk.service.catalog import FunctionInfo
- from databricks.sdk.service.sql import StatementParameterListItem, StatementState
-
-EXECUTE_FUNCTION_ARG_NAME = "__execution_args__"
-DEFAULT_EXECUTE_FUNCTION_ARGS = {
- "wait_timeout": "30s",
- "row_limit": 100,
- "byte_limit": 4096,
-}
-UC_TOOL_CLIENT_EXECUTION_TIMEOUT = "UC_TOOL_CLIENT_EXECUTION_TIMEOUT"
-DEFAULT_UC_TOOL_CLIENT_EXECUTION_TIMEOUT = "120"
-_logger = logging.getLogger(__name__)
-
-
-def is_scalar(function: "FunctionInfo") -> bool:
- from databricks.sdk.service.catalog import ColumnTypeName
-
- return function.data_type != ColumnTypeName.TABLE_TYPE
-
-
-@dataclass
-class ParameterizedStatement:
- statement: str
- parameters: List["StatementParameterListItem"]
-
-
-@dataclass
-class FunctionExecutionResult:
- """
- Result of executing a function.
- We always use a string to present the result value for AI model to consume.
- """
-
- error: Optional[str] = None
- format: Optional[Literal["SCALAR", "CSV"]] = None
- value: Optional[str] = None
- truncated: Optional[bool] = None
-
- def to_json(self) -> str:
- data = {k: v for (k, v) in self.__dict__.items() if v is not None}
- return json.dumps(data)
-
-
-def get_execute_function_sql_stmt(
- function: "FunctionInfo", json_params: Dict[str, Any]
-) -> ParameterizedStatement:
- from databricks.sdk.service.catalog import ColumnTypeName
- from databricks.sdk.service.sql import StatementParameterListItem
-
- parts = []
- output_params = []
- if is_scalar(function):
- # TODO: IDENTIFIER(:function) did not work
- parts.append(f"SELECT {function.full_name}(")
- else:
- parts.append(f"SELECT * FROM {function.full_name}(")
- if function.input_params is None or function.input_params.parameters is None:
- assert not json_params, (
- "Function has no parameters but parameters were provided."
- )
- else:
- args = []
- use_named_args = False
- for p in function.input_params.parameters:
- if p.name not in json_params:
- if p.parameter_default is not None:
- use_named_args = True
- else:
- raise ValueError(
- f"Parameter {p.name} is required but not provided."
- )
- else:
- arg_clause = ""
- if use_named_args:
- arg_clause += f"{p.name} => "
- json_value = json_params[p.name]
- if p.type_name in (
- ColumnTypeName.ARRAY,
- ColumnTypeName.MAP,
- ColumnTypeName.STRUCT,
- ):
- # Use from_json to restore values of complex types.
- json_value_str = json.dumps(json_value)
- # TODO: parametrize type
- arg_clause += f"from_json(:{p.name}, '{p.type_text}')"
- output_params.append(
- StatementParameterListItem(name=p.name, value=json_value_str)
- )
- elif p.type_name == ColumnTypeName.BINARY:
- # Use ubbase64 to restore binary values.
- arg_clause += f"unbase64(:{p.name})"
- output_params.append(
- StatementParameterListItem(name=p.name, value=json_value)
- )
- else:
- arg_clause += f":{p.name}"
- output_params.append(
- StatementParameterListItem(
- name=p.name, value=json_value, type=p.type_text
- )
- )
- args.append(arg_clause)
- parts.append(",".join(args))
- parts.append(")")
- # TODO: check extra params in kwargs
- statement = "".join(parts)
- return ParameterizedStatement(statement=statement, parameters=output_params)
-
-
-def execute_function(
- ws: "WorkspaceClient",
- warehouse_id: str,
- function: "FunctionInfo",
- parameters: Dict[str, Any],
-) -> FunctionExecutionResult:
- """
- Execute a function with the given arguments and return the result.
- """
- try:
- import pandas as pd
- except ImportError as e:
- raise ImportError(
- "Could not import pandas python package. "
- "Please install it with `pip install pandas`."
- ) from e
- from databricks.sdk.service.sql import StatementState
-
- if (
- function.input_params
- and function.input_params.parameters
- and any(
- p.name == EXECUTE_FUNCTION_ARG_NAME
- for p in function.input_params.parameters
- )
- ):
- raise ValueError(
- "Parameter name conflicts with the reserved argument name for executing "
- f"functions: {EXECUTE_FUNCTION_ARG_NAME}. "
- f"Please rename the parameter {EXECUTE_FUNCTION_ARG_NAME}."
- )
-
- # avoid modifying the original dict
- execute_statement_args = {**DEFAULT_EXECUTE_FUNCTION_ARGS}
- allowed_execute_statement_args = inspect.signature(
- ws.statement_execution.execute_statement
- ).parameters
- if not any(
- p.kind in (p.VAR_POSITIONAL, p.VAR_KEYWORD)
- for p in allowed_execute_statement_args.values()
- ):
- invalid_params = set()
- passed_execute_statement_args = parameters.pop(EXECUTE_FUNCTION_ARG_NAME, {})
- for k, v in passed_execute_statement_args.items():
- if k in allowed_execute_statement_args:
- execute_statement_args[k] = v
- else:
- invalid_params.add(k)
- if invalid_params:
- raise ValueError(
- f"Invalid parameters for executing functions: {invalid_params}. "
- f"Allowed parameters are: {allowed_execute_statement_args.keys()}."
- )
-
- # TODO: async so we can run functions in parallel
- parametrized_statement = get_execute_function_sql_stmt(function, parameters)
- response = ws.statement_execution.execute_statement(
- statement=parametrized_statement.statement,
- warehouse_id=warehouse_id,
- parameters=parametrized_statement.parameters,
- **execute_statement_args,
- )
- if response.status and job_pending(response.status.state) and response.statement_id:
- statement_id = response.statement_id
- wait_time = 0
- retry_cnt = 0
- client_execution_timeout = int(
- os.environ.get(
- UC_TOOL_CLIENT_EXECUTION_TIMEOUT,
- DEFAULT_UC_TOOL_CLIENT_EXECUTION_TIMEOUT,
- )
- )
- while wait_time < client_execution_timeout:
- wait = min(2**retry_cnt, client_execution_timeout - wait_time)
- _logger.debug(
- f"Retrying {retry_cnt} time to get statement execution "
- f"status after {wait} seconds."
- )
- time.sleep(wait)
- response = ws.statement_execution.get_statement(statement_id)
- if response.status is None or not job_pending(response.status.state):
- break
- wait_time += wait
- retry_cnt += 1
- if response.status and job_pending(response.status.state):
- return FunctionExecutionResult(
- error=f"Statement execution is still pending after {wait_time} "
- "seconds. Please increase the wait_timeout argument for executing "
- f"the function or increase {UC_TOOL_CLIENT_EXECUTION_TIMEOUT} "
- "environment variable for increasing retrying time, default is "
- f"{DEFAULT_UC_TOOL_CLIENT_EXECUTION_TIMEOUT} seconds."
- )
- assert response.status is not None, f"Statement execution failed: {response}"
- if response.status.state != StatementState.SUCCEEDED:
- error = response.status.error
- assert error is not None, (
- f"Statement execution failed but no error message was provided: {response}"
- )
- return FunctionExecutionResult(error=f"{error.error_code}: {error.message}")
- manifest = response.manifest
- assert manifest is not None
- truncated = manifest.truncated
- result = response.result
- assert result is not None, (
- "Statement execution succeeded but no result was provided."
- )
- data_array = result.data_array
- if is_scalar(function):
- value = None
- if data_array and len(data_array) > 0 and len(data_array[0]) > 0:
- value = str(data_array[0][0])
- return FunctionExecutionResult(
- format="SCALAR", value=value, truncated=truncated
- )
- else:
- schema = manifest.schema
- assert schema is not None and schema.columns is not None, (
- "Statement execution succeeded but no schema was provided."
- )
- columns = [c.name for c in schema.columns]
- if data_array is None:
- data_array = []
- pdf = pd.DataFrame.from_records(data_array, columns=columns)
- csv_buffer = StringIO()
- pdf.to_csv(csv_buffer, index=False)
- return FunctionExecutionResult(
- format="CSV", value=csv_buffer.getvalue(), truncated=truncated
- )
-
-
-def job_pending(state: Optional["StatementState"]) -> bool:
- from databricks.sdk.service.sql import StatementState
-
- return state in (StatementState.PENDING, StatementState.RUNNING)
diff --git a/libs/community/langchain_community/tools/databricks/tool.py b/libs/community/langchain_community/tools/databricks/tool.py
deleted file mode 100644
index 81932cdfdf..0000000000
--- a/libs/community/langchain_community/tools/databricks/tool.py
+++ /dev/null
@@ -1,210 +0,0 @@
-import json
-from datetime import date, datetime
-from decimal import Decimal
-from hashlib import md5
-from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type, Union
-
-from langchain_core._api import deprecated
-from langchain_core.tools import BaseTool, StructuredTool
-from langchain_core.tools.base import BaseToolkit
-from pydantic import BaseModel, Field, create_model
-from typing_extensions import Self
-
-if TYPE_CHECKING:
- from databricks.sdk.service.catalog import FunctionInfo
-
-from pydantic import ConfigDict
-
-from langchain_community.tools.databricks._execution import execute_function
-
-
-def _uc_type_to_pydantic_type(uc_type_json: Union[str, Dict[str, Any]]) -> Type:
- mapping = {
- "long": int,
- "binary": bytes,
- "boolean": bool,
- "date": date,
- "double": float,
- "float": float,
- "integer": int,
- "short": int,
- "string": str,
- "timestamp": datetime,
- "timestamp_ntz": datetime,
- "byte": int,
- }
- if isinstance(uc_type_json, str):
- if uc_type_json in mapping:
- return mapping[uc_type_json]
- else:
- if uc_type_json.startswith("decimal"):
- return Decimal
- elif uc_type_json == "void" or uc_type_json.startswith("interval"):
- raise TypeError(f"Type {uc_type_json} is not supported.")
- else:
- raise TypeError(
- f"Unknown type {uc_type_json}. Try upgrading this package."
- )
- else:
- assert isinstance(uc_type_json, dict)
- tpe = uc_type_json["type"]
- if tpe == "array":
- element_type = _uc_type_to_pydantic_type(uc_type_json["elementType"])
- if uc_type_json["containsNull"]:
- element_type = Optional[element_type] # type: ignore[assignment]
- return List[element_type] # type: ignore[valid-type]
- elif tpe == "map":
- key_type = uc_type_json["keyType"]
- assert key_type == "string", TypeError(
- f"Only support STRING key type for MAP but got {key_type}."
- )
- value_type = _uc_type_to_pydantic_type(uc_type_json["valueType"])
- if uc_type_json["valueContainsNull"]:
- value_type: Type = Optional[value_type] # type: ignore[no-redef]
- return Dict[str, value_type] # type: ignore[valid-type]
- elif tpe == "struct":
- fields = {}
- for field in uc_type_json["fields"]:
- field_type = _uc_type_to_pydantic_type(field["type"])
- if field.get("nullable"):
- field_type = Optional[field_type] # type: ignore[assignment]
- comment = (
- uc_type_json["metadata"].get("comment")
- if "metadata" in uc_type_json
- else None
- )
- fields[field["name"]] = (field_type, Field(..., description=comment))
- uc_type_json_str = json.dumps(uc_type_json, sort_keys=True)
- type_hash = md5(uc_type_json_str.encode()).hexdigest()[:8]
- return create_model(f"Struct_{type_hash}", **fields) # type: ignore[call-overload]
- else:
- raise TypeError(f"Unknown type {uc_type_json}. Try upgrading this package.")
-
-
-def _generate_args_schema(function: "FunctionInfo") -> Type[BaseModel]:
- if function.input_params is None:
- return BaseModel
- params = function.input_params.parameters
- assert params is not None
- fields = {}
- for p in params:
- assert p.type_json is not None
- type_json = json.loads(p.type_json)["type"]
- pydantic_type = _uc_type_to_pydantic_type(type_json)
- description = p.comment
- default: Any = ...
- if p.parameter_default:
- pydantic_type = Optional[pydantic_type] # type: ignore[assignment]
- default = None
- # TODO: Convert default value string to the correct type.
- # We might need to use statement execution API
- # to get the JSON representation of the value.
- default_description = f"(Default: {p.parameter_default})"
- if description:
- description += f" {default_description}"
- else:
- description = default_description
- fields[p.name] = (
- pydantic_type,
- Field(default=default, description=description),
- )
- return create_model( # type: ignore[call-overload]
- f"{function.catalog_name}__{function.schema_name}__{function.name}__params",
- **fields,
- )
-
-
-def _get_tool_name(function: "FunctionInfo") -> str:
- tool_name = f"{function.catalog_name}__{function.schema_name}__{function.name}"[
- -64:
- ]
- return tool_name
-
-
-def _get_default_workspace_client() -> Any:
- try:
- from databricks.sdk import WorkspaceClient
- except ImportError as e:
- raise ImportError(
- "Could not import databricks-sdk python package. "
- "Please install it with `pip install databricks-sdk`."
- ) from e
- return WorkspaceClient()
-
-
-@deprecated(
- since="0.3.18",
- removal="1.0",
- alternative_import="databricks_langchain.uc_ai.UCFunctionToolkit",
-)
-class UCFunctionToolkit(BaseToolkit):
- warehouse_id: str = Field(
- description="The ID of a Databricks SQL Warehouse to execute functions."
- )
-
- workspace_client: Any = Field(
- default_factory=_get_default_workspace_client,
- description="Databricks workspace client.",
- )
-
- tools: Dict[str, BaseTool] = Field(default_factory=dict)
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- )
-
- def include(self, *function_names: str, **kwargs: Any) -> Self:
- """
- Includes UC functions to the toolkit.
-
- Args:
- functions: A list of UC function names in the format
- "catalog_name.schema_name.function_name" or
- "catalog_name.schema_name.*".
- If the function name ends with ".*",
- all functions in the schema will be added.
- kwargs: Extra arguments to pass to StructuredTool, e.g., `return_direct`.
- """
- for name in function_names:
- if name.endswith(".*"):
- catalog_name, schema_name = name[:-2].split(".")
- # TODO: handle pagination, warn and truncate if too many
- functions = self.workspace_client.functions.list(
- catalog_name=catalog_name, schema_name=schema_name
- )
- for f in functions:
- assert f.full_name is not None
- self.include(f.full_name, **kwargs)
- else:
- if name not in self.tools:
- self.tools[name] = self._make_tool(name, **kwargs)
- return self
-
- def _make_tool(self, function_name: str, **kwargs: Any) -> BaseTool:
- function = self.workspace_client.functions.get(function_name)
- name = _get_tool_name(function)
- description = function.comment or ""
- args_schema = _generate_args_schema(function)
-
- def func(*args: Any, **kwargs: Any) -> str:
- # TODO: We expect all named args and ignore args.
- # Non-empty args show up when the function has no parameters.
- args_json = json.loads(json.dumps(kwargs, default=str))
- result = execute_function(
- ws=self.workspace_client,
- warehouse_id=self.warehouse_id,
- function=function,
- parameters=args_json,
- )
- return result.to_json()
-
- return StructuredTool(
- name=name,
- description=description,
- args_schema=args_schema,
- func=func,
- **kwargs,
- )
-
- def get_tools(self) -> List[BaseTool]:
- return list(self.tools.values())
diff --git a/libs/community/langchain_community/tools/dataforseo_api_search/__init__.py b/libs/community/langchain_community/tools/dataforseo_api_search/__init__.py
deleted file mode 100644
index 1e2cd9efe9..0000000000
--- a/libs/community/langchain_community/tools/dataforseo_api_search/__init__.py
+++ /dev/null
@@ -1,9 +0,0 @@
-from langchain_community.tools.dataforseo_api_search.tool import (
- DataForSeoAPISearchResults,
- DataForSeoAPISearchRun,
-)
-
-"""DataForSeo API Toolkit."""
-"""Tool for the DataForSeo SERP API."""
-
-__all__ = ["DataForSeoAPISearchRun", "DataForSeoAPISearchResults"]
diff --git a/libs/community/langchain_community/tools/dataforseo_api_search/tool.py b/libs/community/langchain_community/tools/dataforseo_api_search/tool.py
deleted file mode 100644
index 65bf1ec22c..0000000000
--- a/libs/community/langchain_community/tools/dataforseo_api_search/tool.py
+++ /dev/null
@@ -1,71 +0,0 @@
-"""Tool for the DataForSeo SERP API."""
-
-from typing import Optional
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from langchain_core.tools import BaseTool
-from pydantic import Field
-
-from langchain_community.utilities.dataforseo_api_search import DataForSeoAPIWrapper
-
-
-class DataForSeoAPISearchRun(BaseTool):
- """Tool that queries the DataForSeo Google search API."""
-
- name: str = "dataforseo_api_search"
- description: str = (
- "A robust Google Search API provided by DataForSeo."
- "This tool is handy when you need information about trending topics "
- "or current events."
- )
- api_wrapper: DataForSeoAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return str(self.api_wrapper.run(query))
-
- async def _arun(
- self,
- query: str,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool asynchronously."""
- return (await self.api_wrapper.arun(query)).__str__()
-
-
-class DataForSeoAPISearchResults(BaseTool):
- """Tool that queries the DataForSeo Google Search API
- and get back json."""
-
- name: str = "dataforseo_results_json"
- description: str = (
- "A comprehensive Google Search API provided by DataForSeo."
- "This tool is useful for obtaining real-time data on current events "
- "or popular searches."
- "The input should be a search query and the output is a JSON object "
- "of the query results."
- )
- api_wrapper: DataForSeoAPIWrapper = Field(default_factory=DataForSeoAPIWrapper)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return str(self.api_wrapper.results(query))
-
- async def _arun(
- self,
- query: str,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool asynchronously."""
- return (await self.api_wrapper.aresults(query)).__str__()
diff --git a/libs/community/langchain_community/tools/dataherald/__init__.py b/libs/community/langchain_community/tools/dataherald/__init__.py
deleted file mode 100644
index 74140e97cf..0000000000
--- a/libs/community/langchain_community/tools/dataherald/__init__.py
+++ /dev/null
@@ -1,7 +0,0 @@
-"""Dataherald API toolkit."""
-
-from langchain_community.tools.dataherald.tool import DataheraldTextToSQL
-
-__all__ = [
- "DataheraldTextToSQL",
-]
diff --git a/libs/community/langchain_community/tools/dataherald/tool.py b/libs/community/langchain_community/tools/dataherald/tool.py
deleted file mode 100644
index 2a2546a328..0000000000
--- a/libs/community/langchain_community/tools/dataherald/tool.py
+++ /dev/null
@@ -1,36 +0,0 @@
-"""Tool for the Dataherald Hosted API"""
-
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.utilities.dataherald import DataheraldAPIWrapper
-
-
-class DataheraldTextToSQLInput(BaseModel):
- prompt: str = Field(
- description="Natural language query to be translated to a SQL query."
- )
-
-
-class DataheraldTextToSQL(BaseTool):
- """Tool that queries using the Dataherald SDK."""
-
- name: str = "dataherald"
- description: str = (
- "A wrapper around Dataherald. "
- "Text to SQL. "
- "Input should be a prompt and an existing db_connection_id"
- )
- api_wrapper: DataheraldAPIWrapper
- args_schema: Type[BaseModel] = DataheraldTextToSQLInput
-
- def _run(
- self,
- prompt: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Dataherald tool."""
- return self.api_wrapper.run(prompt)
diff --git a/libs/community/langchain_community/tools/ddg_search/__init__.py b/libs/community/langchain_community/tools/ddg_search/__init__.py
deleted file mode 100644
index 5b7de286b8..0000000000
--- a/libs/community/langchain_community/tools/ddg_search/__init__.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""DuckDuckGo Search API toolkit."""
-
-from langchain_community.tools.ddg_search.tool import DuckDuckGoSearchRun
-
-__all__ = ["DuckDuckGoSearchRun"]
diff --git a/libs/community/langchain_community/tools/ddg_search/tool.py b/libs/community/langchain_community/tools/ddg_search/tool.py
deleted file mode 100644
index 7db9b77da4..0000000000
--- a/libs/community/langchain_community/tools/ddg_search/tool.py
+++ /dev/null
@@ -1,154 +0,0 @@
-"""Tool for the DuckDuckGo search API."""
-
-import json
-import warnings
-from typing import Any, List, Literal, Optional, Type, Union
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.utilities.duckduckgo_search import DuckDuckGoSearchAPIWrapper
-
-
-class DDGInput(BaseModel):
- """Input for the DuckDuckGo search tool."""
-
- query: str = Field(description="search query to look up")
-
-
-class DuckDuckGoSearchRun(BaseTool):
- """DuckDuckGo tool.
-
- Setup:
- Install ``duckduckgo-search`` and ``langchain-community``.
-
- .. code-block:: bash
-
- pip install -U duckduckgo-search langchain-community
-
- Instantiation:
- .. code-block:: python
-
- from langchain_community.tools import DuckDuckGoSearchResults
-
- tool = DuckDuckGoSearchResults()
-
- Invocation with args:
- .. code-block:: python
-
- tool.invoke("Obama")
-
- .. code-block:: python
-
- '[snippet: Users on X have been widely comparing the boost of support felt for Kamala Harris\' campaign to Barack Obama\'s in 2008., title: Surging Support For Kamala Harris Compared To Obama-Era Energy, link: https://www.msn.com/en-us/news/politics/surging-support-for-kamala-harris-compared-to-obama-era-energy/ar-BB1qzdC0, date: 2024-07-24T18:27:01+00:00, source: Newsweek on MSN.com], [snippet: Harris tried to emulate Obama\'s coalition in 2020 and failed. She may have a better shot at reaching young, Black, and Latino voters this time around., title: Harris May Follow Obama\'s Path to the White House After All, link: https://www.msn.com/en-us/news/politics/harris-may-follow-obama-s-path-to-the-white-house-after-all/ar-BB1qv9d4, date: 2024-07-23T22:42:00+00:00, source: Intelligencer on MSN.com], [snippet: The Republican presidential candidate said in an interview on Fox News that he "wouldn\'t be worried" about Michelle Obama running., title: Donald Trump Responds to Michelle Obama Threat, link: https://www.msn.com/en-us/news/politics/donald-trump-responds-to-michelle-obama-threat/ar-BB1qqtu5, date: 2024-07-22T18:26:00+00:00, source: Newsweek on MSN.com], [snippet: H eading into the weekend at his vacation home in Rehoboth Beach, Del., President Biden was reportedly stewing over Barack Obama\'s role in the orchestrated campaign to force him, title: Opinion | Barack Obama Strikes Again, link: https://www.msn.com/en-us/news/politics/opinion-barack-obama-strikes-again/ar-BB1qrfiy, date: 2024-07-22T21:28:00+00:00, source: The Wall Street Journal on MSN.com]'
-
- Invocation with ToolCall:
-
- .. code-block:: python
-
- tool.invoke({"args": {"query":"Obama"}, "id": "1", "name": tool.name, "type": "tool_call"})
-
- .. code-block:: python
-
- ToolMessage(content="[snippet: Biden, Obama and the Clintons Will Speak at the Democratic Convention. The president, two of his predecessors and the party's 2016 nominee are said to be planning speeches at the party's ..., title: Biden, Obama and the Clintons Will Speak at the Democratic Convention ..., link: https://www.nytimes.com/2024/08/12/us/politics/dnc-speakers-biden-obama-clinton.html], [snippet: Barack Obama—with his wife, Michelle—being sworn in as the 44th president of the United States, January 20, 2009. Key events in the life of Barack Obama. Barack Obama (born August 4, 1961, Honolulu, Hawaii, U.S.) is the 44th president of the United States (2009-17) and the first African American to hold the office., title: Barack Obama | Biography, Parents, Education, Presidency, Books ..., link: https://www.britannica.com/biography/Barack-Obama], [snippet: Former President Barack Obama released a letter about President Biden's decision to drop out of the 2024 presidential race. Notably, Obama did not name or endorse Vice President Kamala Harris., title: Read Obama's full statement on Biden dropping out - CBS News, link: https://www.cbsnews.com/news/barack-obama-biden-dropping-out-2024-presidential-race-full-statement/], [snippet: Many of the marquee names in Democratic politics began quickly lining up behind Vice President Kamala Harris on Sunday, but one towering presence in the party held back: Barack Obama. The former ..., title: Why Obama Hasn't Endorsed Harris - The New York Times, link: https://www.nytimes.com/2024/07/21/us/politics/why-obama-hasnt-endorsed-harris.html]", name='duckduckgo_results_json', tool_call_id='1')
- """ # noqa: E501
-
- name: str = "duckduckgo_search"
- description: str = (
- "A wrapper around DuckDuckGo Search. "
- "Useful for when you need to answer questions about current events. "
- "Input should be a search query."
- )
- api_wrapper: DuckDuckGoSearchAPIWrapper = Field(
- default_factory=DuckDuckGoSearchAPIWrapper
- )
- args_schema: Type[BaseModel] = DDGInput
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return self.api_wrapper.run(query)
-
-
-class DuckDuckGoSearchResults(BaseTool):
- """Tool that queries the DuckDuckGo search API and
- returns the results in `output_format`."""
-
- name: str = "duckduckgo_results_json"
- description: str = (
- "A wrapper around Duck Duck Go Search. "
- "Useful for when you need to answer questions about current events. "
- "Input should be a search query."
- )
- max_results: int = Field(alias="num_results", default=4)
- api_wrapper: DuckDuckGoSearchAPIWrapper = Field(
- default_factory=DuckDuckGoSearchAPIWrapper
- )
- backend: str = "text"
- args_schema: Type[BaseModel] = DDGInput
- keys_to_include: Optional[List[str]] = None
- """Which keys from each result to include. If None all keys are included."""
- results_separator: str = ", "
- """Character for separating results."""
- output_format: Literal["string", "json", "list"] = "string"
- """Output format of the search results.
-
- - 'string': Return a concatenated string of the search results.
- - 'json': Return a JSON string of the search results.
- - 'list': Return a list of dictionaries of the search results.
- """
- response_format: Literal["content_and_artifact"] = "content_and_artifact"
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> tuple[Union[List[dict], str], List[dict]]:
- """Use the tool."""
- raw_results = self.api_wrapper.results(
- query, self.max_results, source=self.backend
- )
- results = [
- {
- k: v
- for k, v in d.items()
- if not self.keys_to_include or k in self.keys_to_include
- }
- for d in raw_results
- ]
-
- if self.output_format == "list":
- return results, raw_results
- elif self.output_format == "json":
- return json.dumps(results), raw_results
- elif self.output_format == "string":
- res_strs = [", ".join([f"{k}: {v}" for k, v in d.items()]) for d in results]
- return self.results_separator.join(res_strs), raw_results
- else:
- raise ValueError(
- f"Invalid output_format: {self.output_format}. "
- "Needs to be one of 'string', 'json', 'list'."
- )
-
-
-def DuckDuckGoSearchTool(*args: Any, **kwargs: Any) -> DuckDuckGoSearchRun:
- """
- Deprecated. Use DuckDuckGoSearchRun instead.
-
- Args:
- *args:
- **kwargs:
-
- Returns:
- DuckDuckGoSearchRun
- """
- warnings.warn(
- "DuckDuckGoSearchTool will be deprecated in the future. "
- "Please use DuckDuckGoSearchRun instead.",
- DeprecationWarning,
- )
- return DuckDuckGoSearchRun(*args, **kwargs)
diff --git a/libs/community/langchain_community/tools/e2b_data_analysis/__init__.py b/libs/community/langchain_community/tools/e2b_data_analysis/__init__.py
deleted file mode 100644
index e69de29bb2..0000000000
diff --git a/libs/community/langchain_community/tools/e2b_data_analysis/tool.py b/libs/community/langchain_community/tools/e2b_data_analysis/tool.py
deleted file mode 100644
index 3f952f6dc6..0000000000
--- a/libs/community/langchain_community/tools/e2b_data_analysis/tool.py
+++ /dev/null
@@ -1,243 +0,0 @@
-from __future__ import annotations
-
-import ast
-import json
-import os
-from io import StringIO
-from sys import version_info
-from typing import IO, TYPE_CHECKING, Any, Callable, List, Optional, Type, Union
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManager,
- CallbackManagerForToolRun,
-)
-from langchain_core.tools import BaseTool, Tool
-from pydantic import BaseModel, Field, PrivateAttr
-
-from langchain_community.tools.e2b_data_analysis.unparse import Unparser
-
-if TYPE_CHECKING:
- from e2b import EnvVars
- from e2b.templates.data_analysis import Artifact
-
-base_description = """Evaluates python code in a sandbox environment. \
-The environment is long running and exists across multiple executions. \
-You must send the whole script every time and print your outputs. \
-Script should be pure python code that can be evaluated. \
-It should be in python format NOT markdown. \
-The code should NOT be wrapped in backticks. \
-All python packages including requests, matplotlib, scipy, numpy, pandas, \
-etc are available. Create and display chart using `plt.show()`."""
-
-
-def _unparse(tree: ast.AST) -> str:
- """Unparse the AST."""
- if version_info.minor < 9:
- s = StringIO()
- Unparser(tree, file=s)
- source_code = s.getvalue()
- s.close()
- else:
- source_code = ast.unparse(tree)
- return source_code
-
-
-def add_last_line_print(code: str) -> str:
- """Add print statement to the last line if it's missing.
-
- Sometimes, the LLM-generated code doesn't have `print(variable_name)`, instead the
- LLM tries to print the variable only by writing `variable_name` (as you would in
- REPL, for example).
-
- This methods checks the AST of the generated Python code and adds the print
- statement to the last line if it's missing.
- """
- tree = ast.parse(code)
- node = tree.body[-1]
- if isinstance(node, ast.Expr) and isinstance(node.value, ast.Call):
- if isinstance(node.value.func, ast.Name) and node.value.func.id == "print":
- return _unparse(tree)
-
- if isinstance(node, ast.Expr):
- tree.body[-1] = ast.Expr(
- value=ast.Call(
- func=ast.Name(id="print", ctx=ast.Load()),
- args=[node.value],
- keywords=[],
- )
- )
-
- return _unparse(tree)
-
-
-class UploadedFile(BaseModel):
- """Description of the uploaded path with its remote path."""
-
- name: str
- remote_path: str
- description: str
-
-
-class E2BDataAnalysisToolArguments(BaseModel):
- """Arguments for the E2BDataAnalysisTool."""
-
- python_code: str = Field(
- ...,
- examples=["print('Hello World')"],
- description=(
- "The python script to be evaluated. "
- "The contents will be in main.py. "
- "It should not be in markdown format."
- ),
- )
-
-
-class E2BDataAnalysisTool(BaseTool):
- """Tool for running python code in a sandboxed environment for data analysis."""
-
- name: str = "e2b_data_analysis"
- args_schema: Type[BaseModel] = E2BDataAnalysisToolArguments
- session: Any
- description: str
- _uploaded_files: List[UploadedFile] = PrivateAttr(default_factory=list)
-
- def __init__(
- self,
- api_key: Optional[str] = None,
- cwd: Optional[str] = None,
- env_vars: Optional[EnvVars] = None,
- on_stdout: Optional[Callable[[str], Any]] = None,
- on_stderr: Optional[Callable[[str], Any]] = None,
- on_artifact: Optional[Callable[[Artifact], Any]] = None,
- on_exit: Optional[Callable[[int], Any]] = None,
- **kwargs: Any,
- ):
- try:
- from e2b import DataAnalysis
- except ImportError as e:
- raise ImportError(
- "Unable to import e2b, please install with `pip install e2b`."
- ) from e
-
- # If no API key is provided, E2B will try to read it from the environment
- # variable E2B_API_KEY
- super().__init__(description=base_description, **kwargs)
- self.session = DataAnalysis(
- api_key=api_key,
- cwd=cwd,
- env_vars=env_vars,
- on_stdout=on_stdout,
- on_stderr=on_stderr,
- on_exit=on_exit,
- on_artifact=on_artifact,
- )
-
- def close(self) -> None:
- """Close the cloud sandbox."""
- self._uploaded_files = []
- self.session.close()
-
- @property
- def uploaded_files_description(self) -> str:
- if len(self._uploaded_files) == 0:
- return ""
- lines = ["The following files available in the sandbox:"]
-
- for f in self._uploaded_files:
- if f.description == "":
- lines.append(f"- path: `{f.remote_path}`")
- else:
- lines.append(
- f"- path: `{f.remote_path}` \n description: `{f.description}`"
- )
- return "\n".join(lines)
-
- def _run(
- self,
- python_code: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- callbacks: Optional[CallbackManager] = None,
- ) -> str:
- python_code = add_last_line_print(python_code)
-
- if callbacks is not None:
- on_artifact = getattr(callbacks.metadata, "on_artifact", None)
- else:
- on_artifact = None
-
- stdout, stderr, artifacts = self.session.run_python(
- python_code, on_artifact=on_artifact
- )
-
- out = {
- "stdout": stdout,
- "stderr": stderr,
- "artifacts": list(map(lambda artifact: artifact.name, artifacts)),
- }
- return json.dumps(out)
-
- async def _arun(
- self,
- python_code: str,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- raise NotImplementedError("e2b_data_analysis does not support async")
-
- def run_command(
- self,
- cmd: str,
- ) -> dict:
- """Run shell command in the sandbox."""
- proc = self.session.process.start(cmd)
- output = proc.wait()
- return {
- "stdout": output.stdout,
- "stderr": output.stderr,
- "exit_code": output.exit_code,
- }
-
- def install_python_packages(self, package_names: Union[str, List[str]]) -> None:
- """Install python packages in the sandbox."""
- self.session.install_python_packages(package_names)
-
- def install_system_packages(self, package_names: Union[str, List[str]]) -> None:
- """Install system packages (via apt) in the sandbox."""
- self.session.install_system_packages(package_names)
-
- def download_file(self, remote_path: str) -> bytes:
- """Download file from the sandbox."""
- return self.session.download_file(remote_path)
-
- def upload_file(self, file: IO, description: str) -> UploadedFile:
- """Upload file to the sandbox.
-
- The file is uploaded to the '/home/user/' path."""
- remote_path = self.session.upload_file(file)
-
- f = UploadedFile(
- name=os.path.basename(file.name),
- remote_path=remote_path,
- description=description,
- )
- self._uploaded_files.append(f)
- self.description = self.description + "\n" + self.uploaded_files_description
- return f
-
- def remove_uploaded_file(self, uploaded_file: UploadedFile) -> None:
- """Remove uploaded file from the sandbox."""
- self.session.filesystem.remove(uploaded_file.remote_path)
- self._uploaded_files = [
- f
- for f in self._uploaded_files
- if f.remote_path != uploaded_file.remote_path
- ]
- self.description = self.description + "\n" + self.uploaded_files_description
-
- def as_tool(self) -> Tool: # type: ignore[override]
- return Tool.from_function(
- func=self._run,
- name=self.name,
- description=self.description,
- args_schema=self.args_schema,
- )
diff --git a/libs/community/langchain_community/tools/e2b_data_analysis/unparse.py b/libs/community/langchain_community/tools/e2b_data_analysis/unparse.py
deleted file mode 100644
index 0690cbe906..0000000000
--- a/libs/community/langchain_community/tools/e2b_data_analysis/unparse.py
+++ /dev/null
@@ -1,745 +0,0 @@
-# mypy: disable-error-code=no-untyped-def
-# Because Python >3.9 doesn't support ast.unparse,
-# we copied the unparse functionality from here:
-# https://github.com/python/cpython/blob/3.8/Tools/parser/unparse.py
-"Usage: unparse.py "
-
-import ast
-import io
-import sys
-import tokenize
-
-# Large float and imaginary literals get turned into infinities in the AST.
-# We unparse those infinities to INFSTR.
-INFSTR = "1e" + repr(sys.float_info.max_10_exp + 1)
-
-
-def interleave(inter, f, seq):
- """Call f on each item in seq, calling inter() in between."""
- seq = iter(seq)
- try:
- f(next(seq))
- except StopIteration:
- pass
- else:
- for x in seq:
- inter()
- f(x)
-
-
-class Unparser:
- """Traverse an AST and
- output source code for the abstract syntax; original formatting
- is disregarded."""
-
- def __init__(self, tree, file=sys.stdout):
- """Unparser(tree, file=sys.stdout) -> None.
- Print the source for tree to file."""
- self.f = file
- self._indent = 0
- self.dispatch(tree)
- self.f.flush()
-
- def fill(self, text=""):
- "Indent a piece of text, according to the current indentation level"
- self.f.write("\n" + " " * self._indent + text)
-
- def write(self, text):
- "Append a piece of text to the current line."
- self.f.write(text)
-
- def enter(self):
- "Print ':', and increase the indentation."
- self.write(":")
- self._indent += 1
-
- def leave(self):
- "Decrease the indentation level."
- self._indent -= 1
-
- def dispatch(self, tree):
- "Dispatcher function, dispatching tree type T to method _T."
- if isinstance(tree, list):
- for t in tree:
- self.dispatch(t)
- return
- meth = getattr(self, "_" + tree.__class__.__name__)
- meth(tree)
-
- ############### Unparsing methods ######################
- # There should be one method per concrete grammar type #
- # Constructors should be grouped by sum type. Ideally, #
- # this would follow the order in the grammar, but #
- # currently doesn't. #
- ########################################################
-
- def _Module(self, tree):
- for stmt in tree.body:
- self.dispatch(stmt)
-
- # stmt
- def _Expr(self, tree):
- self.fill()
- self.dispatch(tree.value)
-
- def _NamedExpr(self, tree):
- self.write("(")
- self.dispatch(tree.target)
- self.write(" := ")
- self.dispatch(tree.value)
- self.write(")")
-
- def _Import(self, t):
- self.fill("import ")
- interleave(lambda: self.write(", "), self.dispatch, t.names)
-
- def _ImportFrom(self, t):
- self.fill("from ")
- self.write("." * t.level)
- if t.module:
- self.write(t.module)
- self.write(" import ")
- interleave(lambda: self.write(", "), self.dispatch, t.names)
-
- def _Assign(self, t):
- self.fill()
- for target in t.targets:
- self.dispatch(target)
- self.write(" = ")
- self.dispatch(t.value)
-
- def _AugAssign(self, t):
- self.fill()
- self.dispatch(t.target)
- self.write(" " + self.binop[t.op.__class__.__name__] + "= ")
- self.dispatch(t.value)
-
- def _AnnAssign(self, t):
- self.fill()
- if not t.simple and isinstance(t.target, ast.Name):
- self.write("(")
- self.dispatch(t.target)
- if not t.simple and isinstance(t.target, ast.Name):
- self.write(")")
- self.write(": ")
- self.dispatch(t.annotation)
- if t.value:
- self.write(" = ")
- self.dispatch(t.value)
-
- def _Return(self, t):
- self.fill("return")
- if t.value:
- self.write(" ")
- self.dispatch(t.value)
-
- def _Pass(self, t):
- self.fill("pass")
-
- def _Break(self, t):
- self.fill("break")
-
- def _Continue(self, t):
- self.fill("continue")
-
- def _Delete(self, t):
- self.fill("del ")
- interleave(lambda: self.write(", "), self.dispatch, t.targets)
-
- def _Assert(self, t):
- self.fill("assert ")
- self.dispatch(t.test)
- if t.msg:
- self.write(", ")
- self.dispatch(t.msg)
-
- def _Global(self, t):
- self.fill("global ")
- interleave(lambda: self.write(", "), self.write, t.names)
-
- def _Nonlocal(self, t):
- self.fill("nonlocal ")
- interleave(lambda: self.write(", "), self.write, t.names)
-
- def _Await(self, t):
- self.write("(")
- self.write("await")
- if t.value:
- self.write(" ")
- self.dispatch(t.value)
- self.write(")")
-
- def _Yield(self, t):
- self.write("(")
- self.write("yield")
- if t.value:
- self.write(" ")
- self.dispatch(t.value)
- self.write(")")
-
- def _YieldFrom(self, t):
- self.write("(")
- self.write("yield from")
- if t.value:
- self.write(" ")
- self.dispatch(t.value)
- self.write(")")
-
- def _Raise(self, t):
- self.fill("raise")
- if not t.exc:
- assert not t.cause
- return
- self.write(" ")
- self.dispatch(t.exc)
- if t.cause:
- self.write(" from ")
- self.dispatch(t.cause)
-
- def _Try(self, t):
- self.fill("try")
- self.enter()
- self.dispatch(t.body)
- self.leave()
- for ex in t.handlers:
- self.dispatch(ex)
- if t.orelse:
- self.fill("else")
- self.enter()
- self.dispatch(t.orelse)
- self.leave()
- if t.finalbody:
- self.fill("finally")
- self.enter()
- self.dispatch(t.finalbody)
- self.leave()
-
- def _ExceptHandler(self, t):
- self.fill("except")
- if t.type:
- self.write(" ")
- self.dispatch(t.type)
- if t.name:
- self.write(" as ")
- self.write(t.name)
- self.enter()
- self.dispatch(t.body)
- self.leave()
-
- def _ClassDef(self, t):
- self.write("\n")
- for deco in t.decorator_list:
- self.fill("@")
- self.dispatch(deco)
- self.fill("class " + t.name)
- self.write("(")
- comma = False
- for e in t.bases:
- if comma:
- self.write(", ")
- else:
- comma = True
- self.dispatch(e)
- for e in t.keywords:
- if comma:
- self.write(", ")
- else:
- comma = True
- self.dispatch(e)
- self.write(")")
-
- self.enter()
- self.dispatch(t.body)
- self.leave()
-
- def _FunctionDef(self, t):
- self.__FunctionDef_helper(t, "def")
-
- def _AsyncFunctionDef(self, t):
- self.__FunctionDef_helper(t, "async def")
-
- def __FunctionDef_helper(self, t, fill_suffix):
- self.write("\n")
- for deco in t.decorator_list:
- self.fill("@")
- self.dispatch(deco)
- def_str = fill_suffix + " " + t.name + "("
- self.fill(def_str)
- self.dispatch(t.args)
- self.write(")")
- if t.returns:
- self.write(" -> ")
- self.dispatch(t.returns)
- self.enter()
- self.dispatch(t.body)
- self.leave()
-
- def _For(self, t):
- self.__For_helper("for ", t)
-
- def _AsyncFor(self, t):
- self.__For_helper("async for ", t)
-
- def __For_helper(self, fill, t):
- self.fill(fill)
- self.dispatch(t.target)
- self.write(" in ")
- self.dispatch(t.iter)
- self.enter()
- self.dispatch(t.body)
- self.leave()
- if t.orelse:
- self.fill("else")
- self.enter()
- self.dispatch(t.orelse)
- self.leave()
-
- def _If(self, t):
- self.fill("if ")
- self.dispatch(t.test)
- self.enter()
- self.dispatch(t.body)
- self.leave()
- # collapse nested ifs into equivalent elifs.
- while t.orelse and len(t.orelse) == 1 and isinstance(t.orelse[0], ast.If):
- t = t.orelse[0]
- self.fill("elif ")
- self.dispatch(t.test)
- self.enter()
- self.dispatch(t.body)
- self.leave()
- # final else
- if t.orelse:
- self.fill("else")
- self.enter()
- self.dispatch(t.orelse)
- self.leave()
-
- def _While(self, t):
- self.fill("while ")
- self.dispatch(t.test)
- self.enter()
- self.dispatch(t.body)
- self.leave()
- if t.orelse:
- self.fill("else")
- self.enter()
- self.dispatch(t.orelse)
- self.leave()
-
- def _With(self, t):
- self.fill("with ")
- interleave(lambda: self.write(", "), self.dispatch, t.items)
- self.enter()
- self.dispatch(t.body)
- self.leave()
-
- def _AsyncWith(self, t):
- self.fill("async with ")
- interleave(lambda: self.write(", "), self.dispatch, t.items)
- self.enter()
- self.dispatch(t.body)
- self.leave()
-
- # expr
- def _JoinedStr(self, t):
- self.write("f")
- string = io.StringIO()
- self._fstring_JoinedStr(t, string.write)
- self.write(repr(string.getvalue()))
-
- def _FormattedValue(self, t):
- self.write("f")
- string = io.StringIO()
- self._fstring_FormattedValue(t, string.write)
- self.write(repr(string.getvalue()))
-
- def _fstring_JoinedStr(self, t, write):
- for value in t.values:
- meth = getattr(self, "_fstring_" + type(value).__name__)
- meth(value, write)
-
- def _fstring_Constant(self, t, write):
- assert isinstance(t.value, str)
- value = t.value.replace("{", "{{").replace("}", "}}")
- write(value)
-
- def _fstring_FormattedValue(self, t, write):
- write("{")
- expr = io.StringIO()
- Unparser(t.value, expr)
- expr = expr.getvalue().rstrip("\n")
- if expr.startswith("{"):
- write(" ") # Separate pair of opening brackets as "{ {"
- write(expr)
- if t.conversion != -1:
- conversion = chr(t.conversion)
- assert conversion in "sra"
- write(f"!{conversion}")
- if t.format_spec:
- write(":")
- meth = getattr(self, "_fstring_" + type(t.format_spec).__name__)
- meth(t.format_spec, write)
- write("}")
-
- def _Name(self, t):
- self.write(t.id)
-
- def _write_constant(self, value):
- if isinstance(value, (float, complex)):
- # Substitute overflowing decimal literal for AST infinities.
- self.write(repr(value).replace("inf", INFSTR))
- else:
- self.write(repr(value))
-
- def _Constant(self, t):
- value = t.value
- if isinstance(value, tuple):
- self.write("(")
- if len(value) == 1:
- self._write_constant(value[0])
- self.write(",")
- else:
- interleave(lambda: self.write(", "), self._write_constant, value)
- self.write(")")
- elif value is ...:
- self.write("...")
- else:
- if t.kind == "u":
- self.write("u")
- self._write_constant(t.value)
-
- def _List(self, t):
- self.write("[")
- interleave(lambda: self.write(", "), self.dispatch, t.elts)
- self.write("]")
-
- def _ListComp(self, t):
- self.write("[")
- self.dispatch(t.elt)
- for gen in t.generators:
- self.dispatch(gen)
- self.write("]")
-
- def _GeneratorExp(self, t):
- self.write("(")
- self.dispatch(t.elt)
- for gen in t.generators:
- self.dispatch(gen)
- self.write(")")
-
- def _SetComp(self, t):
- self.write("{")
- self.dispatch(t.elt)
- for gen in t.generators:
- self.dispatch(gen)
- self.write("}")
-
- def _DictComp(self, t):
- self.write("{")
- self.dispatch(t.key)
- self.write(": ")
- self.dispatch(t.value)
- for gen in t.generators:
- self.dispatch(gen)
- self.write("}")
-
- def _comprehension(self, t):
- if t.is_async:
- self.write(" async for ")
- else:
- self.write(" for ")
- self.dispatch(t.target)
- self.write(" in ")
- self.dispatch(t.iter)
- for if_clause in t.ifs:
- self.write(" if ")
- self.dispatch(if_clause)
-
- def _IfExp(self, t):
- self.write("(")
- self.dispatch(t.body)
- self.write(" if ")
- self.dispatch(t.test)
- self.write(" else ")
- self.dispatch(t.orelse)
- self.write(")")
-
- def _Set(self, t):
- assert t.elts # should be at least one element
- self.write("{")
- interleave(lambda: self.write(", "), self.dispatch, t.elts)
- self.write("}")
-
- def _Dict(self, t):
- self.write("{")
-
- def write_key_value_pair(k, v):
- self.dispatch(k)
- self.write(": ")
- self.dispatch(v)
-
- def write_item(item):
- k, v = item
- if k is None:
- # for dictionary unpacking operator in dicts {**{'y': 2}}
- # see PEP 448 for details
- self.write("**")
- self.dispatch(v)
- else:
- write_key_value_pair(k, v)
-
- interleave(lambda: self.write(", "), write_item, zip(t.keys, t.values))
- self.write("}")
-
- def _Tuple(self, t):
- self.write("(")
- if len(t.elts) == 1:
- elt = t.elts[0]
- self.dispatch(elt)
- self.write(",")
- else:
- interleave(lambda: self.write(", "), self.dispatch, t.elts)
- self.write(")")
-
- unop = {"Invert": "~", "Not": "not", "UAdd": "+", "USub": "-"}
-
- def _UnaryOp(self, t):
- self.write("(")
- self.write(self.unop[t.op.__class__.__name__])
- self.write(" ")
- self.dispatch(t.operand)
- self.write(")")
-
- binop = {
- "Add": "+",
- "Sub": "-",
- "Mult": "*",
- "MatMult": "@",
- "Div": "/",
- "Mod": "%",
- "LShift": "<<",
- "RShift": ">>",
- "BitOr": "|",
- "BitXor": "^",
- "BitAnd": "&",
- "FloorDiv": "//",
- "Pow": "**",
- }
-
- def _BinOp(self, t):
- self.write("(")
- self.dispatch(t.left)
- self.write(" " + self.binop[t.op.__class__.__name__] + " ")
- self.dispatch(t.right)
- self.write(")")
-
- cmpops = {
- "Eq": "==",
- "NotEq": "!=",
- "Lt": "<",
- "LtE": "<=",
- "Gt": ">",
- "GtE": ">=",
- "Is": "is",
- "IsNot": "is not",
- "In": "in",
- "NotIn": "not in",
- }
-
- def _Compare(self, t):
- self.write("(")
- self.dispatch(t.left)
- for o, e in zip(t.ops, t.comparators):
- self.write(" " + self.cmpops[o.__class__.__name__] + " ")
- self.dispatch(e)
- self.write(")")
-
- boolops = {ast.And: "and", ast.Or: "or"}
-
- def _BoolOp(self, t):
- self.write("(")
- s = " %s " % self.boolops[t.op.__class__]
- interleave(lambda: self.write(s), self.dispatch, t.values)
- self.write(")")
-
- def _Attribute(self, t):
- self.dispatch(t.value)
- # Special case: 3.__abs__() is a syntax error, so if t.value
- # is an integer literal then we need to either parenthesize
- # it or add an extra space to get 3 .__abs__().
- if isinstance(t.value, ast.Constant) and isinstance(t.value.value, int):
- self.write(" ")
- self.write(".")
- self.write(t.attr)
-
- def _Call(self, t):
- self.dispatch(t.func)
- self.write("(")
- comma = False
- for e in t.args:
- if comma:
- self.write(", ")
- else:
- comma = True
- self.dispatch(e)
- for e in t.keywords:
- if comma:
- self.write(", ")
- else:
- comma = True
- self.dispatch(e)
- self.write(")")
-
- def _Subscript(self, t):
- self.dispatch(t.value)
- self.write("[")
- if (
- isinstance(t.slice, ast.Index)
- and isinstance(t.slice.value, ast.Tuple)
- and t.slice.value.elts
- ):
- if len(t.slice.value.elts) == 1:
- elt = t.slice.value.elts[0]
- self.dispatch(elt)
- self.write(",")
- else:
- interleave(lambda: self.write(", "), self.dispatch, t.slice.value.elts)
- else:
- self.dispatch(t.slice)
- self.write("]")
-
- def _Starred(self, t):
- self.write("*")
- self.dispatch(t.value)
-
- # slice
- def _Ellipsis(self, t):
- self.write("...")
-
- def _Index(self, t):
- self.dispatch(t.value)
-
- def _Slice(self, t):
- if t.lower:
- self.dispatch(t.lower)
- self.write(":")
- if t.upper:
- self.dispatch(t.upper)
- if t.step:
- self.write(":")
- self.dispatch(t.step)
-
- def _ExtSlice(self, t):
- if len(t.dims) == 1:
- elt = t.dims[0]
- self.dispatch(elt)
- self.write(",")
- else:
- interleave(lambda: self.write(", "), self.dispatch, t.dims)
-
- # argument
- def _arg(self, t):
- self.write(t.arg)
- if t.annotation:
- self.write(": ")
- self.dispatch(t.annotation)
-
- # others
- def _arguments(self, t):
- first = True
- # normal arguments
- all_args = t.posonlyargs + t.args
- defaults = [None] * (len(all_args) - len(t.defaults)) + t.defaults
- for index, elements in enumerate(zip(all_args, defaults), 1):
- a, d = elements
- if first:
- first = False
- else:
- self.write(", ")
- self.dispatch(a)
- if d:
- self.write("=")
- self.dispatch(d)
- if index == len(t.posonlyargs):
- self.write(", /")
-
- # varargs, or bare '*' if no varargs but keyword-only arguments present
- if t.vararg or t.kwonlyargs:
- if first:
- first = False
- else:
- self.write(", ")
- self.write("*")
- if t.vararg:
- self.write(t.vararg.arg)
- if t.vararg.annotation:
- self.write(": ")
- self.dispatch(t.vararg.annotation)
-
- # keyword-only arguments
- if t.kwonlyargs:
- for a, d in zip(t.kwonlyargs, t.kw_defaults):
- if first:
- first = False
- else:
- self.write(", ")
- self.dispatch(a)
- if d:
- self.write("=")
- self.dispatch(d)
-
- # kwargs
- if t.kwarg:
- if first:
- first = False
- else:
- self.write(", ")
- self.write("**" + t.kwarg.arg)
- if t.kwarg.annotation:
- self.write(": ")
- self.dispatch(t.kwarg.annotation)
-
- def _keyword(self, t):
- if t.arg is None:
- self.write("**")
- else:
- self.write(t.arg)
- self.write("=")
- self.dispatch(t.value)
-
- def _Lambda(self, t):
- self.write("(")
- self.write("lambda ")
- self.dispatch(t.args)
- self.write(": ")
- self.dispatch(t.body)
- self.write(")")
-
- def _alias(self, t):
- self.write(t.name)
- if t.asname:
- self.write(" as " + t.asname)
-
- def _withitem(self, t):
- self.dispatch(t.context_expr)
- if t.optional_vars:
- self.write(" as ")
- self.dispatch(t.optional_vars)
-
-
-def roundtrip(filename, output=sys.stdout):
- """Parse a file and pretty-print it to output.
-
- The output is formatted as valid Python source code.
-
- Args:
- filename: The name of the file to parse.
- output: The output stream to write to.
- """
- with open(filename, "rb") as pyfile:
- encoding = tokenize.detect_encoding(pyfile.readline)[0]
- with open(filename, "r", encoding=encoding) as pyfile:
- source = pyfile.read()
- tree = compile(source, filename, "exec", ast.PyCF_ONLY_AST)
- Unparser(tree, output)
diff --git a/libs/community/langchain_community/tools/edenai/__init__.py b/libs/community/langchain_community/tools/edenai/__init__.py
deleted file mode 100644
index 42a08aba80..0000000000
--- a/libs/community/langchain_community/tools/edenai/__init__.py
+++ /dev/null
@@ -1,35 +0,0 @@
-"""Edenai Tools."""
-
-from langchain_community.tools.edenai.audio_speech_to_text import (
- EdenAiSpeechToTextTool,
-)
-from langchain_community.tools.edenai.audio_text_to_speech import (
- EdenAiTextToSpeechTool,
-)
-from langchain_community.tools.edenai.edenai_base_tool import EdenaiTool
-from langchain_community.tools.edenai.image_explicitcontent import (
- EdenAiExplicitImageTool,
-)
-from langchain_community.tools.edenai.image_objectdetection import (
- EdenAiObjectDetectionTool,
-)
-from langchain_community.tools.edenai.ocr_identityparser import (
- EdenAiParsingIDTool,
-)
-from langchain_community.tools.edenai.ocr_invoiceparser import (
- EdenAiParsingInvoiceTool,
-)
-from langchain_community.tools.edenai.text_moderation import (
- EdenAiTextModerationTool,
-)
-
-__all__ = [
- "EdenAiExplicitImageTool",
- "EdenAiObjectDetectionTool",
- "EdenAiParsingIDTool",
- "EdenAiParsingInvoiceTool",
- "EdenAiTextToSpeechTool",
- "EdenAiSpeechToTextTool",
- "EdenAiTextModerationTool",
- "EdenaiTool",
-]
diff --git a/libs/community/langchain_community/tools/edenai/audio_speech_to_text.py b/libs/community/langchain_community/tools/edenai/audio_speech_to_text.py
deleted file mode 100644
index ead38e7d19..0000000000
--- a/libs/community/langchain_community/tools/edenai/audio_speech_to_text.py
+++ /dev/null
@@ -1,105 +0,0 @@
-from __future__ import annotations
-
-import json
-import logging
-import time
-from typing import List, Optional, Type
-
-import requests
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field, HttpUrl, validator
-
-from langchain_community.tools.edenai.edenai_base_tool import EdenaiTool
-
-logger = logging.getLogger(__name__)
-
-
-class SpeechToTextInput(BaseModel):
- query: HttpUrl = Field(description="url of the audio to analyze")
-
-
-class EdenAiSpeechToTextTool(EdenaiTool):
- """Tool that queries the Eden AI Speech To Text API.
-
- for api reference check edenai documentation:
- https://app.edenai.run/bricks/speech/asynchronous-speech-to-text.
-
- To use, you should have
- the environment variable ``EDENAI_API_KEY`` set with your API token.
- You can find your token here: https://app.edenai.run/admin/account/settings
- """
-
- name: str = "edenai_speech_to_text"
- description: str = (
- "A wrapper around edenai Services speech to text "
- "Useful for when you have to convert audio to text."
- "Input should be a url to an audio file."
- )
- args_schema: Type[BaseModel] = SpeechToTextInput
- is_async: bool = True
-
- language: Optional[str] = "en"
- speakers: Optional[int]
- profanity_filter: bool = False
- custom_vocabulary: Optional[List[str]]
-
- feature: str = "audio"
- subfeature: str = "speech_to_text_async"
- base_url: str = "https://api.edenai.run/v2/audio/speech_to_text_async/"
-
- @validator("providers")
- def check_only_one_provider_selected(cls, v: List[str]) -> List[str]:
- """
- This tool has no feature to combine providers results.
- Therefore we only allow one provider
- """
- if len(v) > 1:
- raise ValueError(
- "Please select only one provider. "
- "The feature to combine providers results is not available "
- "for this tool."
- )
- return v
-
- def _wait_processing(self, url: str) -> requests.Response:
- for _ in range(10):
- time.sleep(1)
- audio_analysis_result = self._get_edenai(url)
- temp = audio_analysis_result.json()
- if temp["status"] == "finished":
- if temp["results"][self.providers[0]]["error"] is not None:
- raise Exception(
- f"""EdenAI returned an unexpected response
- {temp["results"][self.providers[0]]["error"]}"""
- )
- else:
- return audio_analysis_result
-
- raise Exception("Edenai speech to text job id processing Timed out")
-
- def _parse_response(self, response: dict) -> str:
- return response["public_id"]
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- all_params = {
- "file_url": query,
- "language": self.language,
- "speakers": self.speakers,
- "profanity_filter": self.profanity_filter,
- "custom_vocabulary": self.custom_vocabulary,
- }
-
- # filter so we don't send val to api when val is `None
- query_params = {k: v for k, v in all_params.items() if v is not None}
-
- job_id = self._call_eden_ai(query_params)
- url = self.base_url + job_id
- audio_analysis_result = self._wait_processing(url)
- result = audio_analysis_result.text
- formatted_text = json.loads(result)
- return formatted_text["results"][self.providers[0]]["text"]
diff --git a/libs/community/langchain_community/tools/edenai/audio_text_to_speech.py b/libs/community/langchain_community/tools/edenai/audio_text_to_speech.py
deleted file mode 100644
index d17c0854f9..0000000000
--- a/libs/community/langchain_community/tools/edenai/audio_text_to_speech.py
+++ /dev/null
@@ -1,122 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Any, Dict, List, Literal, Optional, Type
-
-import requests
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field, model_validator, validator
-
-from langchain_community.tools.edenai.edenai_base_tool import EdenaiTool
-
-logger = logging.getLogger(__name__)
-
-
-class TextToSpeechInput(BaseModel):
- query: str = Field(description="text to generate audio from")
-
-
-class EdenAiTextToSpeechTool(EdenaiTool):
- """Tool that queries the Eden AI Text to speech API.
- for api reference check edenai documentation:
- https://docs.edenai.co/reference/audio_text_to_speech_create.
-
- To use, you should have
- the environment variable ``EDENAI_API_KEY`` set with your API token.
- You can find your token here: https://app.edenai.run/admin/account/settings
-
- """
-
- name: str = "edenai_text_to_speech"
- description: str = (
- "A wrapper around edenai Services text to speech."
- "Useful for when you need to convert text to speech."
- """the output is a string representing the URL of the audio file,
- or the path to the downloaded wav file """
- )
- args_schema: Type[BaseModel] = TextToSpeechInput
-
- language: Optional[str] = "en"
- """
- language of the text passed to the model.
- """
-
- # optional params see api documentation for more info
- return_type: Literal["url", "wav"] = "url"
- rate: Optional[int] = None
- pitch: Optional[int] = None
- volume: Optional[int] = None
- audio_format: Optional[str] = None
- sampling_rate: Optional[int] = None
- voice_models: Dict[str, str] = Field(default_factory=dict)
-
- voice: Literal["MALE", "FEMALE"]
- """voice option : 'MALE' or 'FEMALE' """
-
- feature: str = "audio"
- subfeature: str = "text_to_speech"
-
- @validator("providers")
- def check_only_one_provider_selected(cls, v: List[str]) -> List[str]:
- """
- This tool has no feature to combine providers results.
- Therefore we only allow one provider
- """
- if len(v) > 1:
- raise ValueError(
- "Please select only one provider. "
- "The feature to combine providers results is not available "
- "for this tool."
- )
- return v
-
- @model_validator(mode="before")
- @classmethod
- def check_voice_models_key_is_provider_name(cls, values: dict) -> Any:
- for key in values.get("voice_models", {}).keys():
- if key not in values.get("providers", []):
- raise ValueError(
- "voice_model should be formatted like this "
- "{: }"
- )
- return values
-
- def _download_wav(self, url: str, save_path: str) -> None:
- response = requests.get(url)
- if response.status_code == 200:
- with open(save_path, "wb") as f:
- f.write(response.content)
- else:
- raise ValueError("Error while downloading wav file")
-
- def _parse_response(self, response: list) -> str:
- result = response[0]
- if self.return_type == "url":
- return result["audio_resource_url"]
- else:
- self._download_wav(result["audio_resource_url"], "audio.wav")
- return "audio.wav"
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- all_params = {
- "text": query,
- "language": self.language,
- "option": self.voice,
- "return_type": self.return_type,
- "rate": self.rate,
- "pitch": self.pitch,
- "volume": self.volume,
- "audio_format": self.audio_format,
- "sampling_rate": self.sampling_rate,
- "settings": self.voice_models,
- }
-
- # filter so we don't send val to api when val is `None
- query_params = {k: v for k, v in all_params.items() if v is not None}
-
- return self._call_eden_ai(query_params)
diff --git a/libs/community/langchain_community/tools/edenai/edenai_base_tool.py b/libs/community/langchain_community/tools/edenai/edenai_base_tool.py
deleted file mode 100644
index edc5b582e2..0000000000
--- a/libs/community/langchain_community/tools/edenai/edenai_base_tool.py
+++ /dev/null
@@ -1,150 +0,0 @@
-from __future__ import annotations
-
-import logging
-from abc import abstractmethod
-from typing import Any, Dict, List, Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from langchain_core.utils import secret_from_env
-from pydantic import Field, SecretStr
-
-logger = logging.getLogger(__name__)
-
-
-class EdenaiTool(BaseTool):
- """
- the base tool for all the EdenAI Tools .
- you should have
- the environment variable ``EDENAI_API_KEY`` set with your API token.
- You can find your token here: https://app.edenai.run/admin/account/settings
- """
-
- feature: str
- subfeature: str
- edenai_api_key: Optional[SecretStr] = Field(
- default_factory=secret_from_env("EDENAI_API_KEY", default=None)
- )
- is_async: bool = False
-
- providers: List[str]
- """provider to use for the API call."""
-
- @staticmethod
- def get_user_agent() -> str:
- from langchain_community import __version__
-
- return f"langchain/{__version__}"
-
- def _call_eden_ai(self, query_params: Dict[str, Any]) -> str:
- """
- Make an API call to the EdenAI service with the specified query parameters.
-
- Args:
- query_params (dict): The parameters to include in the API call.
-
- Returns:
- requests.Response: The response from the EdenAI API call.
-
- """
- api_key = self.edenai_api_key.get_secret_value() if self.edenai_api_key else ""
- headers = {
- "Authorization": f"Bearer {api_key}",
- "User-Agent": self.get_user_agent(),
- }
-
- url = f"https://api.edenai.run/v2/{self.feature}/{self.subfeature}"
-
- payload = {
- "providers": str(self.providers),
- "response_as_dict": False,
- "attributes_as_list": True,
- "show_original_response": False,
- }
-
- payload.update(query_params)
-
- response = requests.post(url, json=payload, headers=headers)
-
- self._raise_on_error(response)
-
- try:
- return self._parse_response(response.json())
- except Exception as e:
- raise RuntimeError(f"An error occurred while running tool: {e}")
-
- def _raise_on_error(self, response: requests.Response) -> None:
- if response.status_code >= 500:
- raise Exception(f"EdenAI Server: Error {response.status_code}")
- elif response.status_code >= 400:
- raise ValueError(f"EdenAI received an invalid payload: {response.text}")
- elif response.status_code != 200:
- raise Exception(
- f"EdenAI returned an unexpected response with status "
- f"{response.status_code}: {response.text}"
- )
-
- # case where edenai call succeeded but provider returned an error
- # (eg: rate limit, server error, etc.)
- if self.is_async is False:
- # async call are different and only return a job_id,
- # not the provider response directly
- provider_response = response.json()[0]
- if provider_response.get("status") == "fail":
- err_msg = provider_response["error"]["message"]
- raise ValueError(err_msg)
-
- @abstractmethod
- def _run(
- self, query: str, run_manager: Optional[CallbackManagerForToolRun] = None
- ) -> str:
- pass
-
- @abstractmethod
- def _parse_response(self, response: Any) -> str:
- """Take a dict response and condense it's data in a human readable string"""
- pass
-
- def _get_edenai(self, url: str) -> requests.Response:
- headers = {
- "accept": "application/json",
- "authorization": f"Bearer {self.edenai_api_key}",
- "User-Agent": self.get_user_agent(),
- }
-
- response = requests.get(url, headers=headers)
-
- self._raise_on_error(response)
-
- return response
-
- def _parse_json_multilevel(
- self, extracted_data: dict, formatted_list: list, level: int = 0
- ) -> None:
- for section, subsections in extracted_data.items():
- indentation = " " * level
- if isinstance(subsections, str):
- subsections = subsections.replace("\n", ",")
- formatted_list.append(f"{indentation}{section} : {subsections}")
-
- elif isinstance(subsections, list):
- formatted_list.append(f"{indentation}{section} : ")
- self._list_handling(subsections, formatted_list, level + 1)
-
- elif isinstance(subsections, dict):
- formatted_list.append(f"{indentation}{section} : ")
- self._parse_json_multilevel(subsections, formatted_list, level + 1)
-
- def _list_handling(
- self, subsection_list: list, formatted_list: list, level: int
- ) -> None:
- for list_item in subsection_list:
- if isinstance(list_item, dict):
- self._parse_json_multilevel(list_item, formatted_list, level)
-
- elif isinstance(list_item, list):
- self._list_handling(list_item, formatted_list, level + 1)
-
- else:
- formatted_list.append(f"{' ' * level}{list_item}")
diff --git a/libs/community/langchain_community/tools/edenai/image_explicitcontent.py b/libs/community/langchain_community/tools/edenai/image_explicitcontent.py
deleted file mode 100644
index 50f9f24338..0000000000
--- a/libs/community/langchain_community/tools/edenai/image_explicitcontent.py
+++ /dev/null
@@ -1,73 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field, HttpUrl
-
-from langchain_community.tools.edenai.edenai_base_tool import EdenaiTool
-
-logger = logging.getLogger(__name__)
-
-
-class ExplicitImageInput(BaseModel):
- query: HttpUrl = Field(description="url of the image to analyze")
-
-
-class EdenAiExplicitImageTool(EdenaiTool):
- """Tool that queries the Eden AI Explicit image detection.
-
- for api reference check edenai documentation:
- https://docs.edenai.co/reference/image_explicit_content_create.
-
- To use, you should have
- the environment variable ``EDENAI_API_KEY`` set with your API token.
- You can find your token here: https://app.edenai.run/admin/account/settings
-
- """
-
- name: str = "edenai_image_explicit_content_detection"
-
- description: str = (
- "A wrapper around edenai Services Explicit image detection. "
- """Useful for when you have to extract Explicit Content from images.
- it detects adult only content in images,
- that is generally inappropriate for people under
- the age of 18 and includes nudity, sexual activity,
- pornography, violence, gore content, etc."""
- "Input should be the string url of the image ."
- )
- args_schema: Type[BaseModel] = ExplicitImageInput
-
- combine_available: bool = True
- feature: str = "image"
- subfeature: str = "explicit_content"
-
- def _parse_json(self, json_data: dict) -> str:
- result_str = f"nsfw_likelihood: {json_data['nsfw_likelihood']}\n"
- for idx, found_obj in enumerate(json_data["items"]):
- label = found_obj["label"].lower()
- likelihood = found_obj["likelihood"]
- result_str += f"{idx}: {label} likelihood {likelihood},\n"
-
- return result_str[:-2]
-
- def _parse_response(self, json_data: list) -> str:
- if len(json_data) == 1:
- result = self._parse_json(json_data[0])
- else:
- for entry in json_data:
- if entry.get("provider") == "eden-ai":
- result = self._parse_json(entry)
-
- return result
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- query_params = {"file_url": query, "attributes_as_list": False}
- return self._call_eden_ai(query_params)
diff --git a/libs/community/langchain_community/tools/edenai/image_objectdetection.py b/libs/community/langchain_community/tools/edenai/image_objectdetection.py
deleted file mode 100644
index 491f6ec5b3..0000000000
--- a/libs/community/langchain_community/tools/edenai/image_objectdetection.py
+++ /dev/null
@@ -1,87 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field, HttpUrl
-
-from langchain_community.tools.edenai.edenai_base_tool import EdenaiTool
-
-logger = logging.getLogger(__name__)
-
-
-class ObjectDetectionInput(BaseModel):
- query: HttpUrl = Field(description="url of the image to analyze")
-
-
-class EdenAiObjectDetectionTool(EdenaiTool):
- """Tool that queries the Eden AI Object detection API.
-
- for api reference check edenai documentation:
- https://docs.edenai.co/reference/image_object_detection_create.
-
- To use, you should have
- the environment variable ``EDENAI_API_KEY`` set with your API token.
- You can find your token here: https://app.edenai.run/admin/account/settings
-
- """
-
- name: str = "edenai_object_detection"
-
- description: str = (
- "A wrapper around edenai Services Object Detection . "
- """Useful for when you have to do an to identify and locate
- (with bounding boxes) objects in an image """
- "Input should be the string url of the image to identify."
- )
- args_schema: Type[BaseModel] = ObjectDetectionInput
-
- show_positions: bool = False
-
- feature: str = "image"
- subfeature: str = "object_detection"
-
- def _parse_json(self, json_data: dict) -> str:
- result = []
- label_info = []
-
- for found_obj in json_data["items"]:
- label_str = f"{found_obj['label']} - Confidence {found_obj['confidence']}"
- x_min = found_obj.get("x_min")
- x_max = found_obj.get("x_max")
- y_min = found_obj.get("y_min")
- y_max = found_obj.get("y_max")
- if self.show_positions and all(
- [
- x_min,
- x_max,
- y_min,
- y_max,
- ]
- ): # some providers don't return positions
- label_str += f""",at the position x_min: {x_min}, x_max: {x_max},
- y_min: {y_min}, y_max: {y_max}"""
- label_info.append(label_str)
-
- result.append("\n".join(label_info))
- return "\n\n".join(result)
-
- def _parse_response(self, response: list) -> str:
- if len(response) == 1:
- result = self._parse_json(response[0])
- else:
- for entry in response:
- if entry.get("provider") == "eden-ai":
- result = self._parse_json(entry)
-
- return result
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- query_params = {"file_url": query, "attributes_as_list": False}
- return self._call_eden_ai(query_params)
diff --git a/libs/community/langchain_community/tools/edenai/ocr_identityparser.py b/libs/community/langchain_community/tools/edenai/ocr_identityparser.py
deleted file mode 100644
index f227034591..0000000000
--- a/libs/community/langchain_community/tools/edenai/ocr_identityparser.py
+++ /dev/null
@@ -1,75 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field, HttpUrl
-
-from langchain_community.tools.edenai.edenai_base_tool import EdenaiTool
-
-logger = logging.getLogger(__name__)
-
-
-class IDParsingInput(BaseModel):
- query: HttpUrl = Field(description="url of the document to parse")
-
-
-class EdenAiParsingIDTool(EdenaiTool):
- """Tool that queries the Eden AI Identity parsing API.
-
- for api reference check edenai documentation:
- https://docs.edenai.co/reference/ocr_identity_parser_create.
-
- To use, you should have
- the environment variable ``EDENAI_API_KEY`` set with your API token.
- You can find your token here: https://app.edenai.run/admin/account/settings
-
- """
-
- name: str = "edenai_identity_parsing"
-
- description: str = (
- "A wrapper around edenai Services Identity parsing. "
- "Useful for when you have to extract information from an ID Document "
- "Input should be the string url of the document to parse."
- )
- args_schema: Type[BaseModel] = IDParsingInput
-
- feature: str = "ocr"
- subfeature: str = "identity_parser"
-
- language: Optional[str] = None
- """
- language of the text passed to the model.
- """
-
- def _parse_response(self, response: list) -> str:
- formatted_list: list = []
-
- if len(response) == 1:
- self._parse_json_multilevel(
- response[0]["extracted_data"][0], formatted_list
- )
- else:
- for entry in response:
- if entry.get("provider") == "eden-ai":
- self._parse_json_multilevel(
- entry["extracted_data"][0], formatted_list
- )
-
- return "\n".join(formatted_list)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- query_params = {
- "file_url": query,
- "language": self.language,
- "attributes_as_list": False,
- }
-
- return self._call_eden_ai(query_params)
diff --git a/libs/community/langchain_community/tools/edenai/ocr_invoiceparser.py b/libs/community/langchain_community/tools/edenai/ocr_invoiceparser.py
deleted file mode 100644
index d526647678..0000000000
--- a/libs/community/langchain_community/tools/edenai/ocr_invoiceparser.py
+++ /dev/null
@@ -1,78 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field, HttpUrl
-
-from langchain_community.tools.edenai.edenai_base_tool import EdenaiTool
-
-logger = logging.getLogger(__name__)
-
-
-class InvoiceParsingInput(BaseModel):
- query: HttpUrl = Field(description="url of the document to parse")
-
-
-class EdenAiParsingInvoiceTool(EdenaiTool):
- """Tool that queries the Eden AI Invoice parsing API.
-
- for api reference check edenai documentation:
- https://docs.edenai.co/reference/ocr_invoice_parser_create.
-
- To use, you should have
- the environment variable ``EDENAI_API_KEY`` set with your API token.
- You can find your token here: https://app.edenai.run/admin/account/settings
-
- """
-
- name: str = "edenai_invoice_parsing"
- description: str = (
- "A wrapper around edenai Services invoice parsing. "
- """Useful for when you have to extract information from
- an image it enables to take invoices
- in a variety of formats and returns the data in contains
- (items, prices, addresses, vendor name, etc.)
- in a structured format to automate the invoice processing """
- "Input should be the string url of the document to parse."
- )
- args_schema: Type[BaseModel] = InvoiceParsingInput
-
- language: Optional[str] = None
- """
- language of the image passed to the model.
- """
-
- feature: str = "ocr"
- subfeature: str = "invoice_parser"
-
- def _parse_response(self, response: list) -> str:
- formatted_list: list = []
-
- if len(response) == 1:
- self._parse_json_multilevel(
- response[0]["extracted_data"][0], formatted_list
- )
- else:
- for entry in response:
- if entry.get("provider") == "eden-ai":
- self._parse_json_multilevel(
- entry["extracted_data"][0], formatted_list
- )
-
- return "\n".join(formatted_list)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- query_params = {
- "file_url": query,
- "language": self.language,
- "attributes_as_list": False,
- }
-
- return self._call_eden_ai(query_params)
diff --git a/libs/community/langchain_community/tools/edenai/text_moderation.py b/libs/community/langchain_community/tools/edenai/text_moderation.py
deleted file mode 100644
index f5f8497ff3..0000000000
--- a/libs/community/langchain_community/tools/edenai/text_moderation.py
+++ /dev/null
@@ -1,78 +0,0 @@
-from __future__ import annotations
-
-import logging
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.edenai.edenai_base_tool import EdenaiTool
-
-logger = logging.getLogger(__name__)
-
-
-class TextModerationInput(BaseModel):
- query: str = Field(description="Text to moderate")
-
-
-class EdenAiTextModerationTool(EdenaiTool):
- """Tool that queries the Eden AI Explicit text detection.
-
- for api reference check edenai documentation:
- https://docs.edenai.co/reference/image_explicit_content_create.
-
- To use, you should have
- the environment variable ``EDENAI_API_KEY`` set with your API token.
- You can find your token here: https://app.edenai.run/admin/account/settings
-
- """
-
- name: str = "edenai_explicit_content_detection_text"
- description: str = (
- "A wrapper around edenai Services explicit content detection for text. "
- """Useful for when you have to scan text for offensive,
- sexually explicit or suggestive content,
- it checks also if there is any content of self-harm,
- violence, racist or hate speech."""
- """the structure of the output is :
- 'the type of the explicit content : the likelihood of it being explicit'
- the likelihood is a number
- between 1 and 5, 1 being the lowest and 5 the highest.
- something is explicit if the likelihood is equal or higher than 3.
- for example :
- nsfw_likelihood: 1
- this is not explicit.
- for example :
- nsfw_likelihood: 3
- this is explicit.
- """
- "Input should be a string."
- )
- args_schema: Type[BaseModel] = TextModerationInput
-
- language: str
-
- feature: str = "text"
- subfeature: str = "moderation"
-
- def _parse_response(self, response: list) -> str:
- formatted_result = []
- for result in response:
- if "nsfw_likelihood" in result.keys():
- formatted_result.append(
- "nsfw_likelihood: " + str(result["nsfw_likelihood"])
- )
-
- for label, likelihood in zip(result["label"], result["likelihood"]):
- formatted_result.append(f'"{label}": {str(likelihood)}')
-
- return "\n".join(formatted_result)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- query_params = {"text": query, "language": self.language}
- return self._call_eden_ai(query_params)
diff --git a/libs/community/langchain_community/tools/eleven_labs/__init__.py b/libs/community/langchain_community/tools/eleven_labs/__init__.py
deleted file mode 100644
index 3cb16a4160..0000000000
--- a/libs/community/langchain_community/tools/eleven_labs/__init__.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""Eleven Labs Services Tools."""
-
-from langchain_community.tools.eleven_labs.text2speech import ElevenLabsText2SpeechTool
-
-__all__ = ["ElevenLabsText2SpeechTool"]
diff --git a/libs/community/langchain_community/tools/eleven_labs/models.py b/libs/community/langchain_community/tools/eleven_labs/models.py
deleted file mode 100644
index 72e699a781..0000000000
--- a/libs/community/langchain_community/tools/eleven_labs/models.py
+++ /dev/null
@@ -1,9 +0,0 @@
-from enum import Enum
-
-
-class ElevenLabsModel(str, Enum):
- """Models available for Eleven Labs Text2Speech."""
-
- MULTI_LINGUAL = "eleven_multilingual_v2"
- MULTI_LINGUAL_FLASH = "eleven_flash_v2_5"
- MONO_LINGUAL = "eleven_flash_v2"
diff --git a/libs/community/langchain_community/tools/eleven_labs/text2speech.py b/libs/community/langchain_community/tools/eleven_labs/text2speech.py
deleted file mode 100644
index 91fd89b379..0000000000
--- a/libs/community/langchain_community/tools/eleven_labs/text2speech.py
+++ /dev/null
@@ -1,92 +0,0 @@
-import tempfile
-from enum import Enum
-from typing import Any, Dict, Optional, Union
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from langchain_core.utils import get_from_dict_or_env
-from pydantic import model_validator
-
-
-def _import_elevenlabs() -> Any:
- try:
- import elevenlabs
- except ImportError as e:
- raise ImportError(
- "Cannot import elevenlabs, please install `pip install elevenlabs`."
- ) from e
- return elevenlabs
-
-
-class ElevenLabsModel(str, Enum):
- """Models available for Eleven Labs Text2Speech."""
-
- MULTI_LINGUAL = "eleven_multilingual_v2"
- MULTI_LINGUAL_FLASH = "eleven_flash_v2_5"
- MONO_LINGUAL = "eleven_flash_v2"
-
-
-class ElevenLabsText2SpeechTool(BaseTool):
- """Tool that queries the Eleven Labs Text2Speech API.
-
- In order to set this up, follow instructions at:
- https://elevenlabs.io/docs
- """
-
- model: Union[ElevenLabsModel, str] = ElevenLabsModel.MULTI_LINGUAL
- voice: str = "JBFqnCBsd6RMkjVDRZzb"
-
- name: str = "eleven_labs_text2speech"
- description: str = (
- "A wrapper around Eleven Labs Text2Speech. "
- "Useful for when you need to convert text to speech. "
- "It supports more than 30 languages, including English, German, Polish, "
- "Spanish, Italian, French, Portuguese, and Hindi. "
- )
-
- @model_validator(mode="before")
- @classmethod
- def validate_environment(cls, values: Dict) -> Any:
- """Validate that api key exists in environment."""
- _ = get_from_dict_or_env(values, "elevenlabs_api_key", "ELEVENLABS_API_KEY")
-
- return values
-
- def _run(
- self, query: str, run_manager: Optional[CallbackManagerForToolRun] = None
- ) -> str:
- """Use the tool."""
- elevenlabs = _import_elevenlabs()
- client = elevenlabs.client.ElevenLabs()
- try:
- speech = client.text_to_speech.convert(
- text=query,
- model_id=self.model,
- voice_id=self.voice,
- output_format="mp3_44100_128",
- )
- with tempfile.NamedTemporaryFile(
- mode="bx", suffix=".mp3", delete=False
- ) as f:
- f.write(speech)
- return f.name
- except Exception as e:
- raise RuntimeError(f"Error while running ElevenLabsText2SpeechTool: {e}")
-
- def play(self, speech_file: str) -> None:
- """Play the text as speech."""
- elevenlabs = _import_elevenlabs()
- with open(speech_file, mode="rb") as f:
- speech = f.read()
-
- elevenlabs.play(speech)
-
- def stream_speech(self, query: str) -> None:
- """Stream the text as speech as it is generated.
- Play the text in your speakers."""
- elevenlabs = _import_elevenlabs()
- client = elevenlabs.client.ElevenLabs()
- speech_stream = client.text_to_speech.convert_as_stream(
- text=query, model_id=self.model, voice_id=self.voice
- )
- elevenlabs.stream(speech_stream)
diff --git a/libs/community/langchain_community/tools/few_shot/__init__.py b/libs/community/langchain_community/tools/few_shot/__init__.py
deleted file mode 100644
index e19f14575f..0000000000
--- a/libs/community/langchain_community/tools/few_shot/__init__.py
+++ /dev/null
@@ -1,3 +0,0 @@
-from langchain_community.tools.few_shot.tool import FewShotSQLTool
-
-__all__ = ["FewShotSQLTool"]
diff --git a/libs/community/langchain_community/tools/few_shot/tool.py b/libs/community/langchain_community/tools/few_shot/tool.py
deleted file mode 100644
index a61a24e1f5..0000000000
--- a/libs/community/langchain_community/tools/few_shot/tool.py
+++ /dev/null
@@ -1,46 +0,0 @@
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.example_selectors import BaseExampleSelector
-from langchain_core.prompts import FewShotPromptTemplate, PromptTemplate
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, ConfigDict, Field
-
-
-class _FewShotToolInput(BaseModel):
- question: str = Field(
- ..., description="The question for which we want example SQL queries."
- )
-
-
-class FewShotSQLTool(BaseTool):
- """Tool to get example SQL queries related to an input question."""
-
- name: str = "few_shot_sql"
- description: str = "Tool to get example SQL queries related to an input question."
- args_schema: Type[BaseModel] = _FewShotToolInput
-
- example_selector: BaseExampleSelector = Field(exclude=True)
- example_input_key: str = "input"
- example_query_key: str = "query"
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- )
-
- def _run(
- self,
- question: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Execute the query, return the results or an error message."""
- example_prompt = PromptTemplate.from_template(
- f"User input: {self.example_input_key}\nSQL query: {self.example_query_key}"
- )
- prompt = FewShotPromptTemplate(
- example_prompt=example_prompt,
- example_selector=self.example_selector,
- suffix="",
- input_variables=[self.example_input_key],
- )
- return prompt.format(**{self.example_input_key: question})
diff --git a/libs/community/langchain_community/tools/file_management/__init__.py b/libs/community/langchain_community/tools/file_management/__init__.py
deleted file mode 100644
index 395f5d5ea6..0000000000
--- a/libs/community/langchain_community/tools/file_management/__init__.py
+++ /dev/null
@@ -1,19 +0,0 @@
-"""File Management Tools."""
-
-from langchain_community.tools.file_management.copy import CopyFileTool
-from langchain_community.tools.file_management.delete import DeleteFileTool
-from langchain_community.tools.file_management.file_search import FileSearchTool
-from langchain_community.tools.file_management.list_dir import ListDirectoryTool
-from langchain_community.tools.file_management.move import MoveFileTool
-from langchain_community.tools.file_management.read import ReadFileTool
-from langchain_community.tools.file_management.write import WriteFileTool
-
-__all__ = [
- "CopyFileTool",
- "DeleteFileTool",
- "FileSearchTool",
- "MoveFileTool",
- "ReadFileTool",
- "WriteFileTool",
- "ListDirectoryTool",
-]
diff --git a/libs/community/langchain_community/tools/file_management/copy.py b/libs/community/langchain_community/tools/file_management/copy.py
deleted file mode 100644
index 7679e3c43b..0000000000
--- a/libs/community/langchain_community/tools/file_management/copy.py
+++ /dev/null
@@ -1,53 +0,0 @@
-import shutil
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.file_management.utils import (
- INVALID_PATH_TEMPLATE,
- BaseFileToolMixin,
- FileValidationError,
-)
-
-
-class FileCopyInput(BaseModel):
- """Input for CopyFileTool."""
-
- source_path: str = Field(..., description="Path of the file to copy")
- destination_path: str = Field(..., description="Path to save the copied file")
-
-
-class CopyFileTool(BaseFileToolMixin, BaseTool):
- """Tool that copies a file."""
-
- name: str = "copy_file"
- args_schema: Type[BaseModel] = FileCopyInput
- description: str = "Create a copy of a file in a specified location"
-
- def _run(
- self,
- source_path: str,
- destination_path: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- try:
- source_path_ = self.get_relative_path(source_path)
- except FileValidationError:
- return INVALID_PATH_TEMPLATE.format(
- arg_name="source_path", value=source_path
- )
- try:
- destination_path_ = self.get_relative_path(destination_path)
- except FileValidationError:
- return INVALID_PATH_TEMPLATE.format(
- arg_name="destination_path", value=destination_path
- )
- try:
- shutil.copy2(source_path_, destination_path_, follow_symlinks=False)
- return f"File copied successfully from {source_path} to {destination_path}."
- except Exception as e:
- return "Error: " + str(e)
-
- # TODO: Add aiofiles method
diff --git a/libs/community/langchain_community/tools/file_management/delete.py b/libs/community/langchain_community/tools/file_management/delete.py
deleted file mode 100644
index 33f4b70b28..0000000000
--- a/libs/community/langchain_community/tools/file_management/delete.py
+++ /dev/null
@@ -1,45 +0,0 @@
-import os
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.file_management.utils import (
- INVALID_PATH_TEMPLATE,
- BaseFileToolMixin,
- FileValidationError,
-)
-
-
-class FileDeleteInput(BaseModel):
- """Input for DeleteFileTool."""
-
- file_path: str = Field(..., description="Path of the file to delete")
-
-
-class DeleteFileTool(BaseFileToolMixin, BaseTool):
- """Tool that deletes a file."""
-
- name: str = "file_delete"
- args_schema: Type[BaseModel] = FileDeleteInput
- description: str = "Delete a file"
-
- def _run(
- self,
- file_path: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- try:
- file_path_ = self.get_relative_path(file_path)
- except FileValidationError:
- return INVALID_PATH_TEMPLATE.format(arg_name="file_path", value=file_path)
- if not file_path_.exists():
- return f"Error: no such file or directory: {file_path}"
- try:
- os.remove(file_path_)
- return f"File deleted successfully: {file_path}."
- except Exception as e:
- return "Error: " + str(e)
-
- # TODO: Add aiofiles method
diff --git a/libs/community/langchain_community/tools/file_management/file_search.py b/libs/community/langchain_community/tools/file_management/file_search.py
deleted file mode 100644
index a00aee40b4..0000000000
--- a/libs/community/langchain_community/tools/file_management/file_search.py
+++ /dev/null
@@ -1,62 +0,0 @@
-import fnmatch
-import os
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.file_management.utils import (
- INVALID_PATH_TEMPLATE,
- BaseFileToolMixin,
- FileValidationError,
-)
-
-
-class FileSearchInput(BaseModel):
- """Input for FileSearchTool."""
-
- dir_path: str = Field(
- default=".",
- description="Subdirectory to search in.",
- )
- pattern: str = Field(
- ...,
- description="Unix shell regex, where * matches everything.",
- )
-
-
-class FileSearchTool(BaseFileToolMixin, BaseTool):
- """Tool that searches for files in a subdirectory that match a regex pattern."""
-
- name: str = "file_search"
- args_schema: Type[BaseModel] = FileSearchInput
- description: str = (
- "Recursively search for files in a subdirectory that match the regex pattern"
- )
-
- def _run(
- self,
- pattern: str,
- dir_path: str = ".",
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- try:
- dir_path_ = self.get_relative_path(dir_path)
- except FileValidationError:
- return INVALID_PATH_TEMPLATE.format(arg_name="dir_path", value=dir_path)
- matches = []
- try:
- for root, _, filenames in os.walk(dir_path_):
- for filename in fnmatch.filter(filenames, pattern):
- absolute_path = os.path.join(root, filename)
- relative_path = os.path.relpath(absolute_path, dir_path_)
- matches.append(relative_path)
- if matches:
- return "\n".join(matches)
- else:
- return f"No files found for pattern {pattern} in directory {dir_path}"
- except Exception as e:
- return "Error: " + str(e)
-
- # TODO: Add aiofiles method
diff --git a/libs/community/langchain_community/tools/file_management/list_dir.py b/libs/community/langchain_community/tools/file_management/list_dir.py
deleted file mode 100644
index a8bfdc8e3a..0000000000
--- a/libs/community/langchain_community/tools/file_management/list_dir.py
+++ /dev/null
@@ -1,46 +0,0 @@
-import os
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.file_management.utils import (
- INVALID_PATH_TEMPLATE,
- BaseFileToolMixin,
- FileValidationError,
-)
-
-
-class DirectoryListingInput(BaseModel):
- """Input for ListDirectoryTool."""
-
- dir_path: str = Field(default=".", description="Subdirectory to list.")
-
-
-class ListDirectoryTool(BaseFileToolMixin, BaseTool):
- """Tool that lists files and directories in a specified folder."""
-
- name: str = "list_directory"
- args_schema: Type[BaseModel] = DirectoryListingInput
- description: str = "List files and directories in a specified folder"
-
- def _run(
- self,
- dir_path: str = ".",
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- try:
- dir_path_ = self.get_relative_path(dir_path)
- except FileValidationError:
- return INVALID_PATH_TEMPLATE.format(arg_name="dir_path", value=dir_path)
- try:
- entries = os.listdir(dir_path_)
- if entries:
- return "\n".join(entries)
- else:
- return f"No files found in directory {dir_path}"
- except Exception as e:
- return "Error: " + str(e)
-
- # TODO: Add aiofiles method
diff --git a/libs/community/langchain_community/tools/file_management/move.py b/libs/community/langchain_community/tools/file_management/move.py
deleted file mode 100644
index 935625172e..0000000000
--- a/libs/community/langchain_community/tools/file_management/move.py
+++ /dev/null
@@ -1,56 +0,0 @@
-import shutil
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.file_management.utils import (
- INVALID_PATH_TEMPLATE,
- BaseFileToolMixin,
- FileValidationError,
-)
-
-
-class FileMoveInput(BaseModel):
- """Input for MoveFileTool."""
-
- source_path: str = Field(..., description="Path of the file to move")
- destination_path: str = Field(..., description="New path for the moved file")
-
-
-class MoveFileTool(BaseFileToolMixin, BaseTool):
- """Tool that moves a file."""
-
- name: str = "move_file"
- args_schema: Type[BaseModel] = FileMoveInput
- description: str = "Move or rename a file from one location to another"
-
- def _run(
- self,
- source_path: str,
- destination_path: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- try:
- source_path_ = self.get_relative_path(source_path)
- except FileValidationError:
- return INVALID_PATH_TEMPLATE.format(
- arg_name="source_path", value=source_path
- )
- try:
- destination_path_ = self.get_relative_path(destination_path)
- except FileValidationError:
- return INVALID_PATH_TEMPLATE.format(
- arg_name="destination_path_", value=destination_path_
- )
- if not source_path_.exists():
- return f"Error: no such file or directory {source_path}"
- try:
- # shutil.move expects str args in 3.8
- shutil.move(str(source_path_), destination_path_)
- return f"File moved successfully from {source_path} to {destination_path}."
- except Exception as e:
- return "Error: " + str(e)
-
- # TODO: Add aiofiles method
diff --git a/libs/community/langchain_community/tools/file_management/read.py b/libs/community/langchain_community/tools/file_management/read.py
deleted file mode 100644
index 9f746ed16c..0000000000
--- a/libs/community/langchain_community/tools/file_management/read.py
+++ /dev/null
@@ -1,45 +0,0 @@
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.file_management.utils import (
- INVALID_PATH_TEMPLATE,
- BaseFileToolMixin,
- FileValidationError,
-)
-
-
-class ReadFileInput(BaseModel):
- """Input for ReadFileTool."""
-
- file_path: str = Field(..., description="name of file")
-
-
-class ReadFileTool(BaseFileToolMixin, BaseTool):
- """Tool that reads a file."""
-
- name: str = "read_file"
- args_schema: Type[BaseModel] = ReadFileInput
- description: str = "Read file from disk"
-
- def _run(
- self,
- file_path: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- try:
- read_path = self.get_relative_path(file_path)
- except FileValidationError:
- return INVALID_PATH_TEMPLATE.format(arg_name="file_path", value=file_path)
- if not read_path.exists():
- return f"Error: no such file or directory: {file_path}"
- try:
- with read_path.open("r", encoding="utf-8") as f:
- content = f.read()
- return content
- except Exception as e:
- return "Error: " + str(e)
-
- # TODO: Add aiofiles method
diff --git a/libs/community/langchain_community/tools/file_management/utils.py b/libs/community/langchain_community/tools/file_management/utils.py
deleted file mode 100644
index 788823fecd..0000000000
--- a/libs/community/langchain_community/tools/file_management/utils.py
+++ /dev/null
@@ -1,54 +0,0 @@
-import sys
-from pathlib import Path
-from typing import Optional
-
-from pydantic import BaseModel
-
-
-def is_relative_to(path: Path, root: Path) -> bool:
- """Check if path is relative to root."""
- if sys.version_info >= (3, 9):
- # No need for a try/except block in Python 3.8+.
- return path.is_relative_to(root)
- try:
- path.relative_to(root)
- return True
- except ValueError:
- return False
-
-
-INVALID_PATH_TEMPLATE = (
- "Error: Access denied to {arg_name}: {value}."
- " Permission granted exclusively to the current working directory"
-)
-
-
-class FileValidationError(ValueError):
- """Error for paths outside the root directory."""
-
-
-class BaseFileToolMixin(BaseModel):
- """Mixin for file system tools."""
-
- root_dir: Optional[str] = None
- """The final path will be chosen relative to root_dir if specified."""
-
- def get_relative_path(self, file_path: str) -> Path:
- """Get the relative path, returning an error if unsupported."""
- if self.root_dir is None:
- return Path(file_path)
- return get_validated_relative_path(Path(self.root_dir), file_path)
-
-
-def get_validated_relative_path(root: Path, user_path: str) -> Path:
- """Resolve a relative path, raising an error if not within the root directory."""
- # Note, this still permits symlinks from outside that point within the root.
- # Further validation would be needed if those are to be disallowed.
- root = root.resolve()
- full_path = (root / user_path).resolve()
-
- if not is_relative_to(full_path, root):
- raise FileValidationError(
- f"Path {user_path} is outside of the allowed directory {root}"
- )
- return full_path
diff --git a/libs/community/langchain_community/tools/file_management/write.py b/libs/community/langchain_community/tools/file_management/write.py
deleted file mode 100644
index 1d62065bb7..0000000000
--- a/libs/community/langchain_community/tools/file_management/write.py
+++ /dev/null
@@ -1,51 +0,0 @@
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.file_management.utils import (
- INVALID_PATH_TEMPLATE,
- BaseFileToolMixin,
- FileValidationError,
-)
-
-
-class WriteFileInput(BaseModel):
- """Input for WriteFileTool."""
-
- file_path: str = Field(..., description="name of file")
- text: str = Field(..., description="text to write to file")
- append: bool = Field(
- default=False, description="Whether to append to an existing file."
- )
-
-
-class WriteFileTool(BaseFileToolMixin, BaseTool):
- """Tool that writes a file to disk."""
-
- name: str = "write_file"
- args_schema: Type[BaseModel] = WriteFileInput
- description: str = "Write file to disk"
-
- def _run(
- self,
- file_path: str,
- text: str,
- append: bool = False,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- try:
- write_path = self.get_relative_path(file_path)
- except FileValidationError:
- return INVALID_PATH_TEMPLATE.format(arg_name="file_path", value=file_path)
- try:
- write_path.parent.mkdir(exist_ok=True, parents=False)
- mode = "a" if append else "w"
- with write_path.open(mode, encoding="utf-8") as f:
- f.write(text)
- return f"File written successfully to {file_path}."
- except Exception as e:
- return "Error: " + str(e)
-
- # TODO: Add aiofiles method
diff --git a/libs/community/langchain_community/tools/financial_datasets/__init__.py b/libs/community/langchain_community/tools/financial_datasets/__init__.py
deleted file mode 100644
index c9deb30d83..0000000000
--- a/libs/community/langchain_community/tools/financial_datasets/__init__.py
+++ /dev/null
@@ -1,17 +0,0 @@
-"""financial datasets tools."""
-
-from langchain_community.tools.financial_datasets.balance_sheets import (
- BalanceSheets,
-)
-from langchain_community.tools.financial_datasets.cash_flow_statements import (
- CashFlowStatements,
-)
-from langchain_community.tools.financial_datasets.income_statements import (
- IncomeStatements,
-)
-
-__all__ = [
- "BalanceSheets",
- "CashFlowStatements",
- "IncomeStatements",
-]
diff --git a/libs/community/langchain_community/tools/financial_datasets/balance_sheets.py b/libs/community/langchain_community/tools/financial_datasets/balance_sheets.py
deleted file mode 100644
index 21508bc6f9..0000000000
--- a/libs/community/langchain_community/tools/financial_datasets/balance_sheets.py
+++ /dev/null
@@ -1,62 +0,0 @@
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.utilities.financial_datasets import FinancialDatasetsAPIWrapper
-
-
-class BalanceSheetsSchema(BaseModel):
- """Input for BalanceSheets."""
-
- ticker: str = Field(
- description="The ticker symbol to fetch balance sheets for.",
- )
- period: str = Field(
- description="The period of the balance sheets. "
- "Possible values are: "
- "annual, quarterly, ttm. "
- "Default is 'annual'.",
- )
- limit: int = Field(
- description="The number of balance sheets to return. Default is 10.",
- )
-
-
-class BalanceSheets(BaseTool):
- """
- Tool that gets balance sheets for a given ticker over a given period.
- """
-
- mode: str = "get_balance_sheets"
- name: str = "balance_sheets"
- description: str = (
- "A wrapper around financial datasets's Balance Sheets API. "
- "This tool is useful for fetching balance sheets for a given ticker."
- "The tool fetches balance sheets for a given ticker over a given period."
- "The period can be annual, quarterly, or trailing twelve months (ttm)."
- "The number of balance sheets to return can also be "
- "specified using the limit parameter."
- )
- args_schema: Type[BalanceSheetsSchema] = BalanceSheetsSchema
-
- api_wrapper: FinancialDatasetsAPIWrapper = Field(..., exclude=True)
-
- def __init__(self, api_wrapper: FinancialDatasetsAPIWrapper):
- super().__init__(api_wrapper=api_wrapper)
-
- def _run(
- self,
- ticker: str,
- period: str,
- limit: Optional[int],
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Balance Sheets API tool."""
- return self.api_wrapper.run(
- mode=self.mode,
- ticker=ticker,
- period=period,
- limit=limit,
- )
diff --git a/libs/community/langchain_community/tools/financial_datasets/cash_flow_statements.py b/libs/community/langchain_community/tools/financial_datasets/cash_flow_statements.py
deleted file mode 100644
index 065c645420..0000000000
--- a/libs/community/langchain_community/tools/financial_datasets/cash_flow_statements.py
+++ /dev/null
@@ -1,62 +0,0 @@
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.utilities.financial_datasets import FinancialDatasetsAPIWrapper
-
-
-class CashFlowStatementsSchema(BaseModel):
- """Input for CashFlowStatements."""
-
- ticker: str = Field(
- description="The ticker symbol to fetch cash flow statements for.",
- )
- period: str = Field(
- description="The period of the cash flow statement. "
- "Possible values are: "
- "annual, quarterly, ttm. "
- "Default is 'annual'.",
- )
- limit: int = Field(
- description="The number of cash flow statements to return. Default is 10.",
- )
-
-
-class CashFlowStatements(BaseTool):
- """
- Tool that gets cash flow statements for a given ticker over a given period.
- """
-
- mode: str = "get_cash_flow_statements"
- name: str = "cash_flow_statements"
- description: str = (
- "A wrapper around financial datasets's Cash Flow Statements API. "
- "This tool is useful for fetching cash flow statements for a given ticker."
- "The tool fetches cash flow statements for a given ticker over a given period."
- "The period can be annual, quarterly, or trailing twelve months (ttm)."
- "The number of cash flow statements to return can also be "
- "specified using the limit parameter."
- )
- args_schema: Type[CashFlowStatementsSchema] = CashFlowStatementsSchema
-
- api_wrapper: FinancialDatasetsAPIWrapper = Field(..., exclude=True)
-
- def __init__(self, api_wrapper: FinancialDatasetsAPIWrapper):
- super().__init__(api_wrapper=api_wrapper)
-
- def _run(
- self,
- ticker: str,
- period: str,
- limit: Optional[int],
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Cash Flow Statements API tool."""
- return self.api_wrapper.run(
- mode=self.mode,
- ticker=ticker,
- period=period,
- limit=limit,
- )
diff --git a/libs/community/langchain_community/tools/financial_datasets/income_statements.py b/libs/community/langchain_community/tools/financial_datasets/income_statements.py
deleted file mode 100644
index c4801f3d06..0000000000
--- a/libs/community/langchain_community/tools/financial_datasets/income_statements.py
+++ /dev/null
@@ -1,62 +0,0 @@
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.utilities.financial_datasets import FinancialDatasetsAPIWrapper
-
-
-class IncomeStatementsSchema(BaseModel):
- """Input for IncomeStatements."""
-
- ticker: str = Field(
- description="The ticker symbol to fetch income statements for.",
- )
- period: str = Field(
- description="The period of the income statement. "
- "Possible values are: "
- "annual, quarterly, ttm. "
- "Default is 'annual'.",
- )
- limit: int = Field(
- description="The number of income statements to return. Default is 10.",
- )
-
-
-class IncomeStatements(BaseTool):
- """
- Tool that gets income statements for a given ticker over a given period.
- """
-
- mode: str = "get_income_statements"
- name: str = "income_statements"
- description: str = (
- "A wrapper around financial datasets's Income Statements API. "
- "This tool is useful for fetching income statements for a given ticker."
- "The tool fetches income statements for a given ticker over a given period."
- "The period can be annual, quarterly, or trailing twelve months (ttm)."
- "The number of income statements to return can also be "
- "specified using the limit parameter."
- )
- args_schema: Type[IncomeStatementsSchema] = IncomeStatementsSchema
-
- api_wrapper: FinancialDatasetsAPIWrapper = Field(..., exclude=True)
-
- def __init__(self, api_wrapper: FinancialDatasetsAPIWrapper):
- super().__init__(api_wrapper=api_wrapper)
-
- def _run(
- self,
- ticker: str,
- period: str,
- limit: Optional[int],
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Income Statements API tool."""
- return self.api_wrapper.run(
- mode=self.mode,
- ticker=ticker,
- period=period,
- limit=limit,
- )
diff --git a/libs/community/langchain_community/tools/github/__init__.py b/libs/community/langchain_community/tools/github/__init__.py
deleted file mode 100644
index 11c741aa55..0000000000
--- a/libs/community/langchain_community/tools/github/__init__.py
+++ /dev/null
@@ -1 +0,0 @@
-"""GitHub Tool"""
diff --git a/libs/community/langchain_community/tools/github/prompt.py b/libs/community/langchain_community/tools/github/prompt.py
deleted file mode 100644
index c75d750407..0000000000
--- a/libs/community/langchain_community/tools/github/prompt.py
+++ /dev/null
@@ -1,109 +0,0 @@
-# flake8: noqa
-GET_ISSUES_PROMPT = """
-This tool will fetch a list of the repository's issues. It will return the title, and issue number of 5 issues. It takes no input."""
-
-GET_ISSUE_PROMPT = """
-This tool will fetch the title, body, and comment thread of a specific issue. **VERY IMPORTANT**: You must specify the issue number as an integer."""
-
-COMMENT_ON_ISSUE_PROMPT = """
-This tool is useful when you need to comment on a GitHub issue. Simply pass in the issue number and the comment you would like to make. Please use this sparingly as we don't want to clutter the comment threads. **VERY IMPORTANT**: Your input to this tool MUST strictly follow these rules:
-
-- First you must specify the issue number as an integer
-- Then you must place two newlines
-- Then you must specify your comment"""
-
-CREATE_PULL_REQUEST_PROMPT = """
-This tool is useful when you need to create a new pull request in a GitHub repository. **VERY IMPORTANT**: Your input to this tool MUST strictly follow these rules:
-
-- First you must specify the title of the pull request
-- Then you must place two newlines
-- Then you must write the body or description of the pull request
-
-When appropriate, always reference relevant issues in the body by using the syntax `closes #>>> OLD
-- Then you must specify the new contents which you would like to replace the old contents with wrapped in NEW <<<< and >>>> NEW
-
-For example, if you would like to replace the contents of the file /test/test.txt from "old contents" to "new contents", you would pass in the following string:
-
-test/test.txt
-
-This is text that will not be changed
-OLD <<<<
-old contents
->>>> OLD
-NEW <<<<
-new contents
->>>> NEW"""
-
-DELETE_FILE_PROMPT = """
-This tool is a wrapper for the GitHub API, useful when you need to delete a file in a GitHub repository. Simply pass in the full file path of the file you would like to delete. **IMPORTANT**: the path must not start with a slash"""
-
-GET_PR_PROMPT = """
-This tool will fetch the title, body, comment thread and commit history of a specific Pull Request (by PR number). **VERY IMPORTANT**: You must specify the PR number as an integer."""
-
-LIST_PRS_PROMPT = """
-This tool will fetch a list of the repository's Pull Requests (PRs). It will return the title, and PR number of 5 PRs. It takes no input."""
-
-LIST_PULL_REQUEST_FILES = """
-This tool will fetch the full text of all files in a pull request (PR) given the PR number as an input. This is useful for understanding the code changes in a PR or contributing to it. **VERY IMPORTANT**: You must specify the PR number as an integer input parameter."""
-
-OVERVIEW_EXISTING_FILES_IN_MAIN = """
-This tool will provide an overview of all existing files in the main branch of the repository. It will list the file names, their respective paths, and a brief summary of their contents. This can be useful for understanding the structure and content of the repository, especially when navigating through large codebases. No input parameters are required."""
-
-OVERVIEW_EXISTING_FILES_BOT_BRANCH = """
-This tool will provide an overview of all files in your current working branch where you should implement changes. This is great for getting a high level overview of the structure of your code. No input parameters are required."""
-
-SEARCH_ISSUES_AND_PRS_PROMPT = """
-This tool will search for issues and pull requests in the repository. **VERY IMPORTANT**: You must specify the search query as a string input parameter."""
-
-SEARCH_CODE_PROMPT = """
-This tool will search for code in the repository. **VERY IMPORTANT**: You must specify the search query as a string input parameter."""
-
-CREATE_REVIEW_REQUEST_PROMPT = """
-This tool will create a review request on the open pull request that matches the current active branch. **VERY IMPORTANT**: You must specify the username of the person who is being requested as a string input parameter."""
-
-LIST_BRANCHES_IN_REPO_PROMPT = """
-This tool will fetch a list of all branches in the repository. It will return the name of each branch. No input parameters are required."""
-
-SET_ACTIVE_BRANCH_PROMPT = """
-This tool will set the active branch in the repository, similar to `git checkout ` and `git switch -c `. **VERY IMPORTANT**: You must specify the name of the branch as a string input parameter."""
-
-CREATE_BRANCH_PROMPT = """
-This tool will create a new branch in the repository. **VERY IMPORTANT**: You must specify the name of the new branch as a string input parameter."""
-
-GET_FILES_FROM_DIRECTORY_PROMPT = """
-This tool will fetch a list of all files in a specified directory. **VERY IMPORTANT**: You must specify the path of the directory as a string input parameter."""
-
-GET_LATEST_RELEASE_PROMPT = """
-This tool will fetch the latest release of the repository. No input parameters are required."""
-
-GET_RELEASES_PROMPT = """
-This tool will fetch the latest 5 releases of the repository. No input parameters are required."""
-
-GET_RELEASE_PROMPT = """
-This tool will fetch a specific release of the repository. **VERY IMPORTANT**: You must specify the tag name of the release as a string input parameter."""
diff --git a/libs/community/langchain_community/tools/github/tool.py b/libs/community/langchain_community/tools/github/tool.py
deleted file mode 100644
index 836ebc0913..0000000000
--- a/libs/community/langchain_community/tools/github/tool.py
+++ /dev/null
@@ -1,52 +0,0 @@
-"""
-This tool allows agents to interact with the pygithub library
-and operate on a GitHub repository.
-
-To use this tool, you must first set as environment variables:
- GITHUB_API_TOKEN
- GITHUB_REPOSITORY -> format: {owner}/{repo}
-
-"""
-
-from typing import Any, Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.utilities.github import GitHubAPIWrapper
-
-
-class GitHubAction(BaseTool):
- """Tool for interacting with the GitHub API."""
-
- api_wrapper: GitHubAPIWrapper = Field(default_factory=GitHubAPIWrapper)
- mode: str
- name: str = ""
- description: str = ""
- args_schema: Optional[Type[BaseModel]] = None
-
- def _run(
- self,
- instructions: Optional[str] = "",
- run_manager: Optional[CallbackManagerForToolRun] = None,
- **kwargs: Any,
- ) -> str:
- """Use the GitHub API to run an operation."""
- if not instructions or instructions == "{}":
- # Catch other forms of empty input that GPT-4 likes to send.
- instructions = ""
- if self.args_schema is not None:
- field_names = list(self.args_schema.schema()["properties"].keys())
- if len(field_names) > 1:
- raise AssertionError(
- f"Expected one argument in tool schema, got {field_names}."
- )
- if field_names:
- field = field_names[0]
- else:
- field = ""
- query = str(kwargs.get(field, ""))
- else:
- query = instructions
- return self.api_wrapper.run(self.mode, query)
diff --git a/libs/community/langchain_community/tools/gitlab/__init__.py b/libs/community/langchain_community/tools/gitlab/__init__.py
deleted file mode 100644
index 75ad8d0196..0000000000
--- a/libs/community/langchain_community/tools/gitlab/__init__.py
+++ /dev/null
@@ -1 +0,0 @@
-"""GitLab Tool"""
diff --git a/libs/community/langchain_community/tools/gitlab/prompt.py b/libs/community/langchain_community/tools/gitlab/prompt.py
deleted file mode 100644
index e8a33ccb57..0000000000
--- a/libs/community/langchain_community/tools/gitlab/prompt.py
+++ /dev/null
@@ -1,94 +0,0 @@
-# flake8: noqa
-GET_ISSUES_PROMPT = """
-This tool will fetch a list of the repository's issues. It will return the title, and issue number of 5 issues. It takes no input.
-"""
-
-GET_ISSUE_PROMPT = """
-This tool will fetch the title, body, and comment thread of a specific issue. **VERY IMPORTANT**: You must specify the issue number as an integer.
-"""
-
-COMMENT_ON_ISSUE_PROMPT = """
-This tool is useful when you need to comment on a GitLab issue. Simply pass in the issue number and the comment you would like to make. Please use this sparingly as we don't want to clutter the comment threads. **VERY IMPORTANT**: Your input to this tool MUST strictly follow these rules:
-
-- First you must specify the issue number as an integer
-- Then you must place two newlines
-- Then you must specify your comment
-"""
-CREATE_PULL_REQUEST_PROMPT = """
-This tool is useful when you need to create a new pull request in a GitLab repository. **VERY IMPORTANT**: Your input to this tool MUST strictly follow these rules:
-
-- First you must specify the title of the pull request
-- Then you must place two newlines
-- Then you must write the body or description of the pull request
-
-To reference an issue in the body, put its issue number directly after a #.
-For example, if you would like to create a pull request called "README updates" with contents "added contributors' names, closes issue #3", you would pass in the following string:
-
-README updates
-
-added contributors' names, closes issue #3
-"""
-CREATE_FILE_PROMPT = """
-This tool is a wrapper for the GitLab API, useful when you need to create a file in a GitLab repository. **VERY IMPORTANT**: Your input to this tool MUST strictly follow these rules:
-
-- First you must specify which file to create by passing a full file path (**IMPORTANT**: the path must not start with a slash)
-- Then you must specify the contents of the file
-
-For example, if you would like to create a file called /test/test.txt with contents "test contents", you would pass in the following string:
-
-test/test.txt
-
-test contents
-"""
-
-READ_FILE_PROMPT = """
-This tool is a wrapper for the GitLab API, useful when you need to read the contents of a file in a GitLab repository. Simply pass in the full file path of the file you would like to read. **IMPORTANT**: the path must not start with a slash
-"""
-
-UPDATE_FILE_PROMPT = """
-This tool is a wrapper for the GitLab API, useful when you need to update the contents of a file in a GitLab repository. **VERY IMPORTANT**: Your input to this tool MUST strictly follow these rules:
-
-- First you must specify which file to modify by passing a full file path (**IMPORTANT**: the path must not start with a slash)
-- Then you must specify the old contents which you would like to replace wrapped in OLD <<<< and >>>> OLD
-- Then you must specify the new contents which you would like to replace the old contents with wrapped in NEW <<<< and >>>> NEW
-
-For example, if you would like to replace the contents of the file /test/test.txt from "old contents" to "new contents", you would pass in the following string:
-
-test/test.txt
-
-This is text that will not be changed
-OLD <<<<
-old contents
->>>> OLD
-NEW <<<<
-new contents
->>>> NEW
-"""
-
-DELETE_FILE_PROMPT = """
-This tool is a wrapper for the GitLab API, useful when you need to delete a file in a GitLab repository. Simply pass in the full file path of the file you would like to delete. **IMPORTANT**: the path must not start with a slash
-"""
-
-GET_REPO_FILES_IN_MAIN = """
-This tool will provide an overview of all existing files in the main branch of the GitLab repository repository. It will list the file names. No input parameters are required.
-"""
-
-GET_REPO_FILES_IN_BOT_BRANCH = """
-This tool will provide an overview of all files in your current working branch where you should implement changes. No input parameters are required.
-"""
-
-GET_REPO_FILES_FROM_DIRECTORY = """
-This tool will provide an overview of all files in your current working branch from a specific directory. **VERY IMPORTANT**: You must specify the path of the directory as a string input parameter.
-"""
-
-LIST_REPO_BRANCES = """
-This tool is a wrapper for the GitLab API, useful when you need to read the branches names in a GitLab repository. No input parameters are required.
-"""
-
-CREATE_REPO_BRANCH = """
-This tool will create a new branch in the repository. **VERY IMPORTANT**: You must specify the name of the new branch as a string input parameter.
-"""
-
-SET_ACTIVE_BRANCH = """
-This tool will set the active branch in the repository, similar to `git checkout ` and `git switch -c `. **VERY IMPORTANT**: You must specify the name of the branch as a string input parameter.
-"""
diff --git a/libs/community/langchain_community/tools/gitlab/tool.py b/libs/community/langchain_community/tools/gitlab/tool.py
deleted file mode 100644
index 338165ec84..0000000000
--- a/libs/community/langchain_community/tools/gitlab/tool.py
+++ /dev/null
@@ -1,34 +0,0 @@
-"""
-This tool allows agents to interact with the python-gitlab library
-and operate on a GitLab repository.
-
-To use this tool, you must first set as environment variables:
- GITLAB_PRIVATE_ACCESS_TOKEN
- GITLAB_REPOSITORY -> format: {owner}/{repo}
-
-"""
-
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import Field
-
-from langchain_community.utilities.gitlab import GitLabAPIWrapper
-
-
-class GitLabAction(BaseTool):
- """Tool for interacting with the GitLab API."""
-
- api_wrapper: GitLabAPIWrapper = Field(default_factory=GitLabAPIWrapper)
- mode: str
- name: str = ""
- description: str = ""
-
- def _run(
- self,
- instructions: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the GitLab API to run an operation."""
- return self.api_wrapper.run(self.mode, instructions)
diff --git a/libs/community/langchain_community/tools/gmail/__init__.py b/libs/community/langchain_community/tools/gmail/__init__.py
deleted file mode 100644
index 7ef66e21dc..0000000000
--- a/libs/community/langchain_community/tools/gmail/__init__.py
+++ /dev/null
@@ -1,17 +0,0 @@
-"""Gmail tools."""
-
-from langchain_community.tools.gmail.create_draft import GmailCreateDraft
-from langchain_community.tools.gmail.get_message import GmailGetMessage
-from langchain_community.tools.gmail.get_thread import GmailGetThread
-from langchain_community.tools.gmail.search import GmailSearch
-from langchain_community.tools.gmail.send_message import GmailSendMessage
-from langchain_community.tools.gmail.utils import get_gmail_credentials
-
-__all__ = [
- "GmailCreateDraft",
- "GmailSendMessage",
- "GmailSearch",
- "GmailGetMessage",
- "GmailGetThread",
- "get_gmail_credentials",
-]
diff --git a/libs/community/langchain_community/tools/gmail/base.py b/libs/community/langchain_community/tools/gmail/base.py
deleted file mode 100644
index d55b0d30f8..0000000000
--- a/libs/community/langchain_community/tools/gmail/base.py
+++ /dev/null
@@ -1,38 +0,0 @@
-"""Base class for Gmail tools."""
-
-from __future__ import annotations
-
-from typing import TYPE_CHECKING
-
-from langchain_core.tools import BaseTool
-from pydantic import Field
-
-from langchain_community.tools.gmail.utils import build_resource_service
-
-if TYPE_CHECKING:
- # This is for linting and IDE typehints
- from googleapiclient.discovery import Resource
-else:
- try:
- # We do this so pydantic can resolve the types when instantiating
- from googleapiclient.discovery import Resource
- except ImportError:
- pass
-
-
-class GmailBaseTool(BaseTool):
- """Base class for Gmail tools."""
-
- api_resource: Resource = Field(default_factory=build_resource_service)
-
- @classmethod
- def from_api_resource(cls, api_resource: Resource) -> "GmailBaseTool":
- """Create a tool from an api resource.
-
- Args:
- api_resource: The api resource to use.
-
- Returns:
- A tool.
- """
- return cls(service=api_resource) # type: ignore[call-arg]
diff --git a/libs/community/langchain_community/tools/gmail/create_draft.py b/libs/community/langchain_community/tools/gmail/create_draft.py
deleted file mode 100644
index ec2495aaa4..0000000000
--- a/libs/community/langchain_community/tools/gmail/create_draft.py
+++ /dev/null
@@ -1,87 +0,0 @@
-import base64
-from email.message import EmailMessage
-from typing import List, Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.gmail.base import GmailBaseTool
-
-
-class CreateDraftSchema(BaseModel):
- """Input for CreateDraftTool."""
-
- message: str = Field(
- ...,
- description="The message to include in the draft.",
- )
- to: List[str] = Field(
- ...,
- description="The list of recipients.",
- )
- subject: str = Field(
- ...,
- description="The subject of the message.",
- )
- cc: Optional[List[str]] = Field(
- None,
- description="The list of CC recipients.",
- )
- bcc: Optional[List[str]] = Field(
- None,
- description="The list of BCC recipients.",
- )
-
-
-class GmailCreateDraft(GmailBaseTool):
- """Tool that creates a draft email for Gmail."""
-
- name: str = "create_gmail_draft"
- description: str = (
- "Use this tool to create a draft email with the provided message fields."
- )
- args_schema: Type[CreateDraftSchema] = CreateDraftSchema
-
- def _prepare_draft_message(
- self,
- message: str,
- to: List[str],
- subject: str,
- cc: Optional[List[str]] = None,
- bcc: Optional[List[str]] = None,
- ) -> dict:
- draft_message = EmailMessage()
- draft_message.set_content(message)
-
- draft_message["To"] = ", ".join(to)
- draft_message["Subject"] = subject
- if cc is not None:
- draft_message["Cc"] = ", ".join(cc)
-
- if bcc is not None:
- draft_message["Bcc"] = ", ".join(bcc)
-
- encoded_message = base64.urlsafe_b64encode(draft_message.as_bytes()).decode()
- return {"message": {"raw": encoded_message}}
-
- def _run(
- self,
- message: str,
- to: List[str],
- subject: str,
- cc: Optional[List[str]] = None,
- bcc: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- try:
- create_message = self._prepare_draft_message(message, to, subject, cc, bcc)
- draft = (
- self.api_resource.users()
- .drafts()
- .create(userId="me", body=create_message)
- .execute()
- )
- output = f"Draft created. Draft Id: {draft['id']}"
- return output
- except Exception as e:
- raise Exception(f"An error occurred: {e}")
diff --git a/libs/community/langchain_community/tools/gmail/get_message.py b/libs/community/langchain_community/tools/gmail/get_message.py
deleted file mode 100644
index 6155cb499f..0000000000
--- a/libs/community/langchain_community/tools/gmail/get_message.py
+++ /dev/null
@@ -1,70 +0,0 @@
-import base64
-import email
-from typing import Dict, Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.gmail.base import GmailBaseTool
-from langchain_community.tools.gmail.utils import clean_email_body
-
-
-class SearchArgsSchema(BaseModel):
- """Input for GetMessageTool."""
-
- message_id: str = Field(
- ...,
- description="The unique ID of the email message, retrieved from a search.",
- )
-
-
-class GmailGetMessage(GmailBaseTool):
- """Tool that gets a message by ID from Gmail."""
-
- name: str = "get_gmail_message"
- description: str = (
- "Use this tool to fetch an email by message ID."
- " Returns the thread ID, snippet, body, subject, and sender."
- )
- args_schema: Type[SearchArgsSchema] = SearchArgsSchema
-
- def _run(
- self,
- message_id: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> Dict:
- """Run the tool."""
- query = (
- self.api_resource.users()
- .messages()
- .get(userId="me", format="raw", id=message_id)
- )
- message_data = query.execute()
- raw_message = base64.urlsafe_b64decode(message_data["raw"])
-
- email_msg = email.message_from_bytes(raw_message)
-
- subject = email_msg["Subject"]
- sender = email_msg["From"]
-
- message_body = ""
- if email_msg.is_multipart():
- for part in email_msg.walk():
- ctype = part.get_content_type()
- cdispo = str(part.get("Content-Disposition"))
- if ctype == "text/plain" and "attachment" not in cdispo:
- message_body = part.get_payload(decode=True).decode("utf-8") # type: ignore[union-attr]
- break
- else:
- message_body = email_msg.get_payload(decode=True).decode("utf-8") # type: ignore[union-attr]
-
- body = clean_email_body(message_body)
-
- return {
- "id": message_id,
- "threadId": message_data["threadId"],
- "snippet": message_data["snippet"],
- "body": body,
- "subject": subject,
- "sender": sender,
- }
diff --git a/libs/community/langchain_community/tools/gmail/get_thread.py b/libs/community/langchain_community/tools/gmail/get_thread.py
deleted file mode 100644
index 5e61bd8bb9..0000000000
--- a/libs/community/langchain_community/tools/gmail/get_thread.py
+++ /dev/null
@@ -1,48 +0,0 @@
-from typing import Dict, Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.gmail.base import GmailBaseTool
-
-
-class GetThreadSchema(BaseModel):
- """Input for GetMessageTool."""
-
- # From https://support.google.com/mail/answer/7190?hl=en
- thread_id: str = Field(
- ...,
- description="The thread ID.",
- )
-
-
-class GmailGetThread(GmailBaseTool):
- """Tool that gets a thread by ID from Gmail."""
-
- name: str = "get_gmail_thread"
- description: str = (
- "Use this tool to search for email messages."
- " The input must be a valid Gmail query."
- " The output is a JSON list of messages."
- )
- args_schema: Type[GetThreadSchema] = GetThreadSchema
-
- def _run(
- self,
- thread_id: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> Dict:
- """Run the tool."""
- query = self.api_resource.users().threads().get(userId="me", id=thread_id)
- thread_data = query.execute()
- if not isinstance(thread_data, dict):
- raise ValueError("The output of the query must be a list.")
- messages = thread_data["messages"]
- thread_data["messages"] = []
- keys_to_keep = ["id", "snippet", "snippet"]
- # TODO: Parse body.
- for message in messages:
- thread_data["messages"].append(
- {k: message[k] for k in keys_to_keep if k in message}
- )
- return thread_data
diff --git a/libs/community/langchain_community/tools/gmail/search.py b/libs/community/langchain_community/tools/gmail/search.py
deleted file mode 100644
index eb61968429..0000000000
--- a/libs/community/langchain_community/tools/gmail/search.py
+++ /dev/null
@@ -1,149 +0,0 @@
-import base64
-import email
-from enum import Enum
-from typing import Any, Dict, List, Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.gmail.base import GmailBaseTool
-from langchain_community.tools.gmail.utils import clean_email_body
-
-
-class Resource(str, Enum):
- """Enumerator of Resources to search."""
-
- THREADS = "threads"
- MESSAGES = "messages"
-
-
-class SearchArgsSchema(BaseModel):
- """Input for SearchGmailTool."""
-
- # From https://support.google.com/mail/answer/7190?hl=en
- query: str = Field(
- ...,
- description="The Gmail query. Example filters include from:sender,"
- " to:recipient, subject:subject, -filtered_term,"
- " in:folder, is:important|read|starred, after:year/mo/date, "
- "before:year/mo/date, label:label_name"
- ' "exact phrase".'
- " Search newer/older than using d (day), m (month), and y (year): "
- "newer_than:2d, older_than:1y."
- " Attachments with extension example: filename:pdf. Multiple term"
- " matching example: from:amy OR from:david.",
- )
- resource: Resource = Field(
- default=Resource.MESSAGES,
- description="Whether to search for threads or messages.",
- )
- max_results: int = Field(
- default=10,
- description="The maximum number of results to return.",
- )
-
-
-class GmailSearch(GmailBaseTool):
- """Tool that searches for messages or threads in Gmail."""
-
- name: str = "search_gmail"
- description: str = (
- "Use this tool to search for email messages or threads."
- " The input must be a valid Gmail query."
- " The output is a JSON list of the requested resource."
- )
- args_schema: Type[SearchArgsSchema] = SearchArgsSchema
-
- def _parse_threads(self, threads: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
- # Add the thread message snippets to the thread results
- results = []
- for thread in threads:
- thread_id = thread["id"]
- thread_data = (
- self.api_resource.users()
- .threads()
- .get(userId="me", id=thread_id)
- .execute()
- )
- messages = thread_data["messages"]
- thread["messages"] = []
- for message in messages:
- snippet = message["snippet"]
- thread["messages"].append({"snippet": snippet, "id": message["id"]})
- results.append(thread)
-
- return results
-
- def _parse_messages(self, messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
- results = []
- for message in messages:
- message_id = message["id"]
- message_data = (
- self.api_resource.users()
- .messages()
- .get(userId="me", format="raw", id=message_id)
- .execute()
- )
-
- raw_message = base64.urlsafe_b64decode(message_data["raw"])
-
- email_msg = email.message_from_bytes(raw_message)
-
- subject = email_msg["Subject"]
- sender = email_msg["From"]
-
- message_body = ""
- if email_msg.is_multipart():
- for part in email_msg.walk():
- ctype = part.get_content_type()
- cdispo = str(part.get("Content-Disposition"))
- if ctype == "text/plain" and "attachment" not in cdispo:
- try:
- message_body = part.get_payload(decode=True).decode("utf-8") # type: ignore[union-attr]
- except UnicodeDecodeError:
- message_body = part.get_payload(decode=True).decode( # type: ignore[union-attr]
- "latin-1"
- )
- break
- else:
- message_body = email_msg.get_payload(decode=True).decode("utf-8") # type: ignore[union-attr]
-
- body = clean_email_body(message_body)
-
- results.append(
- {
- "id": message["id"],
- "threadId": message_data["threadId"],
- "snippet": message_data["snippet"],
- "body": body,
- "subject": subject,
- "sender": sender,
- "from": email_msg["From"],
- "date": email_msg["Date"],
- "to": email_msg["To"],
- "cc": email_msg["Cc"],
- }
- )
- return results
-
- def _run(
- self,
- query: str,
- resource: Resource = Resource.MESSAGES,
- max_results: int = 10,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> List[Dict[str, Any]]:
- """Run the tool."""
- results = (
- self.api_resource.users()
- .messages()
- .list(userId="me", q=query, maxResults=max_results)
- .execute()
- .get(resource.value, [])
- )
- if resource == Resource.THREADS:
- return self._parse_threads(results)
- elif resource == Resource.MESSAGES:
- return self._parse_messages(results)
- else:
- raise NotImplementedError(f"Resource of type {resource} not implemented.")
diff --git a/libs/community/langchain_community/tools/gmail/send_message.py b/libs/community/langchain_community/tools/gmail/send_message.py
deleted file mode 100644
index 0d9fbc6697..0000000000
--- a/libs/community/langchain_community/tools/gmail/send_message.py
+++ /dev/null
@@ -1,91 +0,0 @@
-"""Send Gmail messages."""
-
-import base64
-from email.mime.multipart import MIMEMultipart
-from email.mime.text import MIMEText
-from typing import Any, Dict, List, Optional, Type, Union
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.gmail.base import GmailBaseTool
-
-
-class SendMessageSchema(BaseModel):
- """Input for SendMessageTool."""
-
- message: str = Field(
- ...,
- description="The message to send.",
- )
- to: Union[str, List[str]] = Field(
- ...,
- description="The list of recipients.",
- )
- subject: str = Field(
- ...,
- description="The subject of the message.",
- )
- cc: Optional[Union[str, List[str]]] = Field(
- None,
- description="The list of CC recipients.",
- )
- bcc: Optional[Union[str, List[str]]] = Field(
- None,
- description="The list of BCC recipients.",
- )
-
-
-class GmailSendMessage(GmailBaseTool):
- """Tool that sends a message to Gmail."""
-
- name: str = "send_gmail_message"
- description: str = (
- "Use this tool to send email messages. The input is the message, recipients"
- )
- args_schema: Type[SendMessageSchema] = SendMessageSchema
-
- def _prepare_message(
- self,
- message: str,
- to: Union[str, List[str]],
- subject: str,
- cc: Optional[Union[str, List[str]]] = None,
- bcc: Optional[Union[str, List[str]]] = None,
- ) -> Dict[str, Any]:
- """Create a message for an email."""
- mime_message = MIMEMultipart()
- mime_message.attach(MIMEText(message, "html"))
-
- mime_message["To"] = ", ".join(to if isinstance(to, list) else [to])
- mime_message["Subject"] = subject
- if cc is not None:
- mime_message["Cc"] = ", ".join(cc if isinstance(cc, list) else [cc])
-
- if bcc is not None:
- mime_message["Bcc"] = ", ".join(bcc if isinstance(bcc, list) else [bcc])
-
- encoded_message = base64.urlsafe_b64encode(mime_message.as_bytes()).decode()
- return {"raw": encoded_message}
-
- def _run(
- self,
- message: str,
- to: Union[str, List[str]],
- subject: str,
- cc: Optional[Union[str, List[str]]] = None,
- bcc: Optional[Union[str, List[str]]] = None,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Run the tool."""
- try:
- create_message = self._prepare_message(message, to, subject, cc=cc, bcc=bcc)
- send_message = (
- self.api_resource.users()
- .messages()
- .send(userId="me", body=create_message)
- )
- sent_message = send_message.execute()
- return f"Message sent. Message Id: {sent_message['id']}"
- except Exception as error:
- raise Exception(f"An error occurred: {error}")
diff --git a/libs/community/langchain_community/tools/gmail/utils.py b/libs/community/langchain_community/tools/gmail/utils.py
deleted file mode 100644
index e53a453836..0000000000
--- a/libs/community/langchain_community/tools/gmail/utils.py
+++ /dev/null
@@ -1,124 +0,0 @@
-"""Gmail tool utils."""
-
-from __future__ import annotations
-
-import logging
-import os
-from typing import TYPE_CHECKING, List, Optional, Tuple
-
-from langchain_core.utils import guard_import
-
-if TYPE_CHECKING:
- from google.auth.transport.requests import Request
- from google.oauth2.credentials import Credentials
- from google_auth_oauthlib.flow import InstalledAppFlow
- from googleapiclient.discovery import Resource
- from googleapiclient.discovery import build as build_resource
-
-logger = logging.getLogger(__name__)
-
-
-def import_google() -> Tuple[Request, Credentials]:
- """Import google libraries.
-
- Returns:
- Tuple[Request, Credentials]: Request and Credentials classes.
- """
- return (
- guard_import(
- module_name="google.auth.transport.requests",
- pip_name="google-auth-httplib2",
- ).Request,
- guard_import(
- module_name="google.oauth2.credentials", pip_name="google-auth-httplib2"
- ).Credentials,
- )
-
-
-def import_installed_app_flow() -> InstalledAppFlow:
- """Import InstalledAppFlow class.
-
- Returns:
- InstalledAppFlow: InstalledAppFlow class.
- """
- return guard_import(
- module_name="google_auth_oauthlib.flow", pip_name="google-auth-oauthlib"
- ).InstalledAppFlow
-
-
-def import_googleapiclient_resource_builder() -> build_resource:
- """Import googleapiclient.discovery.build function.
-
- Returns:
- build_resource: googleapiclient.discovery.build function.
- """
- return guard_import(
- module_name="googleapiclient.discovery", pip_name="google-api-python-client"
- ).build
-
-
-DEFAULT_SCOPES = ["https://mail.google.com/"]
-DEFAULT_CREDS_TOKEN_FILE = "token.json"
-DEFAULT_CLIENT_SECRETS_FILE = "credentials.json"
-
-
-def get_gmail_credentials(
- token_file: Optional[str] = None,
- client_secrets_file: Optional[str] = None,
- scopes: Optional[List[str]] = None,
-) -> Credentials:
- """Get credentials."""
- # From https://developers.google.com/gmail/api/quickstart/python
- Request, Credentials = import_google()
- InstalledAppFlow = import_installed_app_flow()
- creds = None
- scopes = scopes or DEFAULT_SCOPES
- token_file = token_file or DEFAULT_CREDS_TOKEN_FILE
- client_secrets_file = client_secrets_file or DEFAULT_CLIENT_SECRETS_FILE
- # The file token.json stores the user's access and refresh tokens, and is
- # created automatically when the authorization flow completes for the first
- # time.
- if os.path.exists(token_file):
- creds = Credentials.from_authorized_user_file(token_file, scopes)
- # If there are no (valid) credentials available, let the user log in.
- if not creds or not creds.valid:
- if creds and creds.expired and creds.refresh_token:
- creds.refresh(Request())
- else:
- # https://developers.google.com/gmail/api/quickstart/python#authorize_credentials_for_a_desktop_application # noqa
- flow = InstalledAppFlow.from_client_secrets_file(
- client_secrets_file, scopes
- )
- creds = flow.run_local_server(port=0, open_browser=False)
- # Save the credentials for the next run
- with open(token_file, "w") as token:
- token.write(creds.to_json())
- return creds
-
-
-def build_resource_service(
- credentials: Optional[Credentials] = None,
- service_name: str = "gmail",
- service_version: str = "v1",
-) -> Resource:
- """Build a Gmail service."""
- credentials = credentials or get_gmail_credentials()
- builder = import_googleapiclient_resource_builder()
- return builder(service_name, service_version, credentials=credentials)
-
-
-def clean_email_body(body: str) -> str:
- """Clean email body."""
- try:
- from bs4 import BeautifulSoup
-
- try:
- soup = BeautifulSoup(str(body), "html.parser")
- body = soup.get_text()
- return str(body)
- except Exception as e:
- logger.error(e)
- return str(body)
- except ImportError:
- logger.warning("BeautifulSoup not installed. Skipping cleaning.")
- return str(body)
diff --git a/libs/community/langchain_community/tools/golden_query/__init__.py b/libs/community/langchain_community/tools/golden_query/__init__.py
deleted file mode 100644
index 4c5ae17a13..0000000000
--- a/libs/community/langchain_community/tools/golden_query/__init__.py
+++ /dev/null
@@ -1,7 +0,0 @@
-"""Golden API toolkit."""
-
-from langchain_community.tools.golden_query.tool import GoldenQueryRun
-
-__all__ = [
- "GoldenQueryRun",
-]
diff --git a/libs/community/langchain_community/tools/golden_query/tool.py b/libs/community/langchain_community/tools/golden_query/tool.py
deleted file mode 100644
index 7cc5c72234..0000000000
--- a/libs/community/langchain_community/tools/golden_query/tool.py
+++ /dev/null
@@ -1,34 +0,0 @@
-"""Tool for the Golden API."""
-
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-
-from langchain_community.utilities.golden_query import GoldenQueryAPIWrapper
-
-
-class GoldenQueryRun(BaseTool):
- """Tool that adds the capability to query using the Golden API and get back JSON."""
-
- name: str = "golden_query"
- description: str = (
- "A wrapper around Golden Query API."
- " Useful for getting entities that match"
- " a natural language query from Golden's Knowledge Base."
- "\nExample queries:"
- "\n- companies in nanotech"
- "\n- list of cloud providers starting in 2019"
- "\nInput should be the natural language query."
- "\nOutput is a paginated list of results or an error object"
- " in JSON format."
- )
- api_wrapper: GoldenQueryAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Golden tool."""
- return self.api_wrapper.run(query)
diff --git a/libs/community/langchain_community/tools/google_books.py b/libs/community/langchain_community/tools/google_books.py
deleted file mode 100644
index 572dd2747a..0000000000
--- a/libs/community/langchain_community/tools/google_books.py
+++ /dev/null
@@ -1,38 +0,0 @@
-"""Tool for the Google Books API."""
-
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.utilities.google_books import GoogleBooksAPIWrapper
-
-
-class GoogleBooksQueryInput(BaseModel):
- """Input for the GoogleBooksQuery tool."""
-
- query: str = Field(description="query to look up on google books")
-
-
-class GoogleBooksQueryRun(BaseTool):
- """Tool that searches the Google Books API."""
-
- name: str = "GoogleBooks"
- description: str = (
- "A wrapper around Google Books. "
- "Useful for when you need to answer general inquiries about "
- "books of certain topics and generate recommendation based "
- "off of key words"
- "Input should be a query string"
- )
- api_wrapper: GoogleBooksAPIWrapper
- args_schema: Type[BaseModel] = GoogleBooksQueryInput
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Google Books tool."""
- return self.api_wrapper.run(query)
diff --git a/libs/community/langchain_community/tools/google_cloud/__init__.py b/libs/community/langchain_community/tools/google_cloud/__init__.py
deleted file mode 100644
index ec7deb8951..0000000000
--- a/libs/community/langchain_community/tools/google_cloud/__init__.py
+++ /dev/null
@@ -1,7 +0,0 @@
-"""Google Cloud Tools."""
-
-from langchain_community.tools.google_cloud.texttospeech import (
- GoogleCloudTextToSpeechTool,
-)
-
-__all__ = ["GoogleCloudTextToSpeechTool"]
diff --git a/libs/community/langchain_community/tools/google_cloud/texttospeech.py b/libs/community/langchain_community/tools/google_cloud/texttospeech.py
deleted file mode 100644
index 02a24e9cf1..0000000000
--- a/libs/community/langchain_community/tools/google_cloud/texttospeech.py
+++ /dev/null
@@ -1,97 +0,0 @@
-from __future__ import annotations
-
-import tempfile
-from typing import TYPE_CHECKING, Any, Optional
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-
-from langchain_community.utilities.vertexai import get_client_info
-
-if TYPE_CHECKING:
- from google.cloud import texttospeech
-
-
-def _import_google_cloud_texttospeech() -> Any:
- try:
- from google.cloud import texttospeech
- except ImportError as e:
- raise ImportError(
- "Cannot import google.cloud.texttospeech, please install "
- "`pip install google-cloud-texttospeech`."
- ) from e
- return texttospeech
-
-
-def _encoding_file_extension_map(encoding: texttospeech.AudioEncoding) -> Optional[str]:
- texttospeech = _import_google_cloud_texttospeech()
-
- ENCODING_FILE_EXTENSION_MAP = {
- texttospeech.AudioEncoding.LINEAR16: ".wav",
- texttospeech.AudioEncoding.MP3: ".mp3",
- texttospeech.AudioEncoding.OGG_OPUS: ".ogg",
- texttospeech.AudioEncoding.MULAW: ".wav",
- texttospeech.AudioEncoding.ALAW: ".wav",
- }
- return ENCODING_FILE_EXTENSION_MAP.get(encoding)
-
-
-@deprecated(
- since="0.0.33",
- removal="1.0",
- alternative_import="langchain_google_community.TextToSpeechTool",
-)
-class GoogleCloudTextToSpeechTool(BaseTool):
- """Tool that queries the Google Cloud Text to Speech API.
-
- In order to set this up, follow instructions at:
- https://cloud.google.com/text-to-speech/docs/before-you-begin
- """
-
- name: str = "google_cloud_texttospeech"
- description: str = (
- "A wrapper around Google Cloud Text-to-Speech. "
- "Useful for when you need to synthesize audio from text. "
- "It supports multiple languages, including English, German, Polish, "
- "Spanish, Italian, French, Portuguese, and Hindi. "
- )
-
- _client: Any
-
- def __init__(self, **kwargs: Any) -> None:
- """Initializes private fields."""
- texttospeech = _import_google_cloud_texttospeech()
-
- super().__init__(**kwargs)
-
- self._client = texttospeech.TextToSpeechClient(
- client_info=get_client_info(module="text-to-speech")
- )
-
- def _run(
- self,
- input_text: str,
- language_code: str = "en-US",
- ssml_gender: Optional[texttospeech.SsmlVoiceGender] = None,
- audio_encoding: Optional[texttospeech.AudioEncoding] = None,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- texttospeech = _import_google_cloud_texttospeech()
- ssml_gender = ssml_gender or texttospeech.SsmlVoiceGender.NEUTRAL
- audio_encoding = audio_encoding or texttospeech.AudioEncoding.MP3
-
- response = self._client.synthesize_speech(
- input=texttospeech.SynthesisInput(text=input_text),
- voice=texttospeech.VoiceSelectionParams(
- language_code=language_code, ssml_gender=ssml_gender
- ),
- audio_config=texttospeech.AudioConfig(audio_encoding=audio_encoding),
- )
-
- suffix = _encoding_file_extension_map(audio_encoding)
-
- with tempfile.NamedTemporaryFile(mode="bx", suffix=suffix, delete=False) as f:
- f.write(response.audio_content)
- return f.name
diff --git a/libs/community/langchain_community/tools/google_finance/__init__.py b/libs/community/langchain_community/tools/google_finance/__init__.py
deleted file mode 100644
index bc06ae46d5..0000000000
--- a/libs/community/langchain_community/tools/google_finance/__init__.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""Google Finance API Toolkit."""
-
-from langchain_community.tools.google_finance.tool import GoogleFinanceQueryRun
-
-__all__ = ["GoogleFinanceQueryRun"]
diff --git a/libs/community/langchain_community/tools/google_finance/tool.py b/libs/community/langchain_community/tools/google_finance/tool.py
deleted file mode 100644
index 82eb82de31..0000000000
--- a/libs/community/langchain_community/tools/google_finance/tool.py
+++ /dev/null
@@ -1,29 +0,0 @@
-"""Tool for the Google Finance"""
-
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-
-from langchain_community.utilities.google_finance import GoogleFinanceAPIWrapper
-
-
-class GoogleFinanceQueryRun(BaseTool):
- """Tool that queries the Google Finance API."""
-
- name: str = "google_finance"
- description: str = (
- "A wrapper around Google Finance Search. "
- "Useful for when you need to get information about"
- "google search Finance from Google Finance"
- "Input should be a search query."
- )
- api_wrapper: GoogleFinanceAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return self.api_wrapper.run(query)
diff --git a/libs/community/langchain_community/tools/google_jobs/__init__.py b/libs/community/langchain_community/tools/google_jobs/__init__.py
deleted file mode 100644
index f23e0eecff..0000000000
--- a/libs/community/langchain_community/tools/google_jobs/__init__.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""Google Jobs API Toolkit."""
-
-from langchain_community.tools.google_jobs.tool import GoogleJobsQueryRun
-
-__all__ = ["GoogleJobsQueryRun"]
diff --git a/libs/community/langchain_community/tools/google_jobs/tool.py b/libs/community/langchain_community/tools/google_jobs/tool.py
deleted file mode 100644
index 6a83b3043d..0000000000
--- a/libs/community/langchain_community/tools/google_jobs/tool.py
+++ /dev/null
@@ -1,29 +0,0 @@
-"""Tool for the Google Trends"""
-
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-
-from langchain_community.utilities.google_jobs import GoogleJobsAPIWrapper
-
-
-class GoogleJobsQueryRun(BaseTool):
- """Tool that queries the Google Jobs API."""
-
- name: str = "google_jobs"
- description: str = (
- "A wrapper around Google Jobs Search. "
- "Useful for when you need to get information about"
- "google search Jobs from Google Jobs"
- "Input should be a search query."
- )
- api_wrapper: GoogleJobsAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return self.api_wrapper.run(query)
diff --git a/libs/community/langchain_community/tools/google_lens/__init__.py b/libs/community/langchain_community/tools/google_lens/__init__.py
deleted file mode 100644
index 15a0c17937..0000000000
--- a/libs/community/langchain_community/tools/google_lens/__init__.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""Google Lens API Toolkit."""
-
-from langchain_community.tools.google_lens.tool import GoogleLensQueryRun
-
-__all__ = ["GoogleLensQueryRun"]
diff --git a/libs/community/langchain_community/tools/google_lens/tool.py b/libs/community/langchain_community/tools/google_lens/tool.py
deleted file mode 100644
index 38a4b847e2..0000000000
--- a/libs/community/langchain_community/tools/google_lens/tool.py
+++ /dev/null
@@ -1,29 +0,0 @@
-"""Tool for the Google Lens"""
-
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-
-from langchain_community.utilities.google_lens import GoogleLensAPIWrapper
-
-
-class GoogleLensQueryRun(BaseTool):
- """Tool that queries the Google Lens API."""
-
- name: str = "google_lens"
- description: str = (
- "A wrapper around Google Lens Search. "
- "Useful for when you need to get information related"
- "to an image from Google Lens"
- "Input should be a url to an image."
- )
- api_wrapper: GoogleLensAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return self.api_wrapper.run(query)
diff --git a/libs/community/langchain_community/tools/google_places/__init__.py b/libs/community/langchain_community/tools/google_places/__init__.py
deleted file mode 100644
index 6d3b948ea5..0000000000
--- a/libs/community/langchain_community/tools/google_places/__init__.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""Google Places API Toolkit."""
-
-from langchain_community.tools.google_places.tool import GooglePlacesTool
-
-__all__ = ["GooglePlacesTool"]
diff --git a/libs/community/langchain_community/tools/google_places/tool.py b/libs/community/langchain_community/tools/google_places/tool.py
deleted file mode 100644
index 77a1469073..0000000000
--- a/libs/community/langchain_community/tools/google_places/tool.py
+++ /dev/null
@@ -1,43 +0,0 @@
-"""Tool for the Google search API."""
-
-from typing import Optional, Type
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.utilities.google_places_api import GooglePlacesAPIWrapper
-
-
-class GooglePlacesSchema(BaseModel):
- """Input for GooglePlacesTool."""
-
- query: str = Field(..., description="Query for google maps")
-
-
-@deprecated(
- since="0.0.33",
- removal="1.0",
- alternative_import="langchain_google_community.GooglePlacesTool",
-)
-class GooglePlacesTool(BaseTool):
- """Tool that queries the Google places API."""
-
- name: str = "google_places"
- description: str = (
- "A wrapper around Google Places. "
- "Useful for when you need to validate or "
- "discover addressed from ambiguous text. "
- "Input should be a search query."
- )
- api_wrapper: GooglePlacesAPIWrapper = Field(default_factory=GooglePlacesAPIWrapper)
- args_schema: Type[BaseModel] = GooglePlacesSchema
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return self.api_wrapper.run(query)
diff --git a/libs/community/langchain_community/tools/google_scholar/__init__.py b/libs/community/langchain_community/tools/google_scholar/__init__.py
deleted file mode 100644
index b83e5dfc1e..0000000000
--- a/libs/community/langchain_community/tools/google_scholar/__init__.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""Google Scholar API Toolkit."""
-
-from langchain_community.tools.google_scholar.tool import GoogleScholarQueryRun
-
-__all__ = ["GoogleScholarQueryRun"]
diff --git a/libs/community/langchain_community/tools/google_scholar/tool.py b/libs/community/langchain_community/tools/google_scholar/tool.py
deleted file mode 100644
index 49f8769696..0000000000
--- a/libs/community/langchain_community/tools/google_scholar/tool.py
+++ /dev/null
@@ -1,29 +0,0 @@
-"""Tool for the Google Scholar"""
-
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-
-from langchain_community.utilities.google_scholar import GoogleScholarAPIWrapper
-
-
-class GoogleScholarQueryRun(BaseTool):
- """Tool that queries the Google search API."""
-
- name: str = "google_scholar"
- description: str = (
- "A wrapper around Google Scholar Search. "
- "Useful for when you need to get information about"
- "research papers from Google Scholar"
- "Input should be a search query."
- )
- api_wrapper: GoogleScholarAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return self.api_wrapper.run(query)
diff --git a/libs/community/langchain_community/tools/google_search/__init__.py b/libs/community/langchain_community/tools/google_search/__init__.py
deleted file mode 100644
index 08eccf0a31..0000000000
--- a/libs/community/langchain_community/tools/google_search/__init__.py
+++ /dev/null
@@ -1,8 +0,0 @@
-"""Google Search API Toolkit."""
-
-from langchain_community.tools.google_search.tool import (
- GoogleSearchResults,
- GoogleSearchRun,
-)
-
-__all__ = ["GoogleSearchRun", "GoogleSearchResults"]
diff --git a/libs/community/langchain_community/tools/google_search/tool.py b/libs/community/langchain_community/tools/google_search/tool.py
deleted file mode 100644
index 3ba05079df..0000000000
--- a/libs/community/langchain_community/tools/google_search/tool.py
+++ /dev/null
@@ -1,60 +0,0 @@
-"""Tool for the Google search API."""
-
-from typing import Optional
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-
-from langchain_community.utilities.google_search import GoogleSearchAPIWrapper
-
-
-@deprecated(
- since="0.0.33",
- removal="1.0",
- alternative_import="langchain_google_community.GoogleSearchRun",
-)
-class GoogleSearchRun(BaseTool):
- """Tool that queries the Google search API."""
-
- name: str = "google_search"
- description: str = (
- "A wrapper around Google Search. "
- "Useful for when you need to answer questions about current events. "
- "Input should be a search query."
- )
- api_wrapper: GoogleSearchAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return self.api_wrapper.run(query)
-
-
-@deprecated(
- since="0.0.33",
- removal="1.0",
- alternative_import="langchain_google_community.GoogleSearchResults",
-)
-class GoogleSearchResults(BaseTool):
- """Tool that queries the Google Search API and gets back json."""
-
- name: str = "google_search_results_json"
- description: str = (
- "A wrapper around Google Search. "
- "Useful for when you need to answer questions about current events. "
- "Input should be a search query. Output is a JSON array of the query results"
- )
- num_results: int = 4
- api_wrapper: GoogleSearchAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return str(self.api_wrapper.results(query, self.num_results))
diff --git a/libs/community/langchain_community/tools/google_serper/__init__.py b/libs/community/langchain_community/tools/google_serper/__init__.py
deleted file mode 100644
index 413481a645..0000000000
--- a/libs/community/langchain_community/tools/google_serper/__init__.py
+++ /dev/null
@@ -1,9 +0,0 @@
-from langchain_community.tools.google_serper.tool import (
- GoogleSerperResults,
- GoogleSerperRun,
-)
-
-"""Google Serper API Toolkit."""
-"""Tool for the Serer.dev Google Search API."""
-
-__all__ = ["GoogleSerperRun", "GoogleSerperResults"]
diff --git a/libs/community/langchain_community/tools/google_serper/tool.py b/libs/community/langchain_community/tools/google_serper/tool.py
deleted file mode 100644
index 562dd012a1..0000000000
--- a/libs/community/langchain_community/tools/google_serper/tool.py
+++ /dev/null
@@ -1,70 +0,0 @@
-"""Tool for the Serper.dev Google Search API."""
-
-from typing import Optional
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from langchain_core.tools import BaseTool
-from pydantic import Field
-
-from langchain_community.utilities.google_serper import GoogleSerperAPIWrapper
-
-
-class GoogleSerperRun(BaseTool):
- """Tool that queries the Serper.dev Google search API."""
-
- name: str = "google_serper"
- description: str = (
- "A low-cost Google Search API."
- "Useful for when you need to answer questions about current events."
- "Input should be a search query."
- )
- api_wrapper: GoogleSerperAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return str(self.api_wrapper.run(query))
-
- async def _arun(
- self,
- query: str,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool asynchronously."""
- return (await self.api_wrapper.arun(query)).__str__()
-
-
-class GoogleSerperResults(BaseTool):
- """Tool that queries the Serper.dev Google Search API
- and get back json."""
-
- name: str = "google_serper_results_json"
- description: str = (
- "A low-cost Google Search API."
- "Useful for when you need to answer questions about current events."
- "Input should be a search query. Output is a JSON object of the query results"
- )
- api_wrapper: GoogleSerperAPIWrapper = Field(default_factory=GoogleSerperAPIWrapper)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return str(self.api_wrapper.results(query))
-
- async def _arun(
- self,
- query: str,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool asynchronously."""
-
- return (await self.api_wrapper.aresults(query)).__str__()
diff --git a/libs/community/langchain_community/tools/google_trends/__init__.py b/libs/community/langchain_community/tools/google_trends/__init__.py
deleted file mode 100644
index ca3d58fc59..0000000000
--- a/libs/community/langchain_community/tools/google_trends/__init__.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""Google Trends API Toolkit."""
-
-from langchain_community.tools.google_trends.tool import GoogleTrendsQueryRun
-
-__all__ = ["GoogleTrendsQueryRun"]
diff --git a/libs/community/langchain_community/tools/google_trends/tool.py b/libs/community/langchain_community/tools/google_trends/tool.py
deleted file mode 100644
index 8b2b5dd8bf..0000000000
--- a/libs/community/langchain_community/tools/google_trends/tool.py
+++ /dev/null
@@ -1,29 +0,0 @@
-"""Tool for the Google Trends"""
-
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-
-from langchain_community.utilities.google_trends import GoogleTrendsAPIWrapper
-
-
-class GoogleTrendsQueryRun(BaseTool):
- """Tool that queries the Google trends API."""
-
- name: str = "google_trends"
- description: str = (
- "A wrapper around Google Trends Search. "
- "Useful for when you need to get information about"
- "google search trends from Google Trends"
- "Input should be a search query."
- )
- api_wrapper: GoogleTrendsAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return self.api_wrapper.run(query)
diff --git a/libs/community/langchain_community/tools/graphql/__init__.py b/libs/community/langchain_community/tools/graphql/__init__.py
deleted file mode 100644
index 7e9a84c377..0000000000
--- a/libs/community/langchain_community/tools/graphql/__init__.py
+++ /dev/null
@@ -1 +0,0 @@
-"""Tools for interacting with a GraphQL API"""
diff --git a/libs/community/langchain_community/tools/graphql/tool.py b/libs/community/langchain_community/tools/graphql/tool.py
deleted file mode 100644
index 0530f8cae0..0000000000
--- a/libs/community/langchain_community/tools/graphql/tool.py
+++ /dev/null
@@ -1,36 +0,0 @@
-import json
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import ConfigDict
-
-from langchain_community.utilities.graphql import GraphQLAPIWrapper
-
-
-class BaseGraphQLTool(BaseTool):
- """Base tool for querying a GraphQL API."""
-
- graphql_wrapper: GraphQLAPIWrapper
-
- name: str = "query_graphql"
- description: str = """\
- Input to this tool is a detailed and correct GraphQL query, output is a result from the API.
- If the query is not correct, an error message will be returned.
- If an error is returned with 'Bad request' in it, rewrite the query and try again.
- If an error is returned with 'Unauthorized' in it, do not try again, but tell the user to change their authentication.
-
- Example Input: query {{ allUsers {{ id, name, email }} }}\
- """ # noqa: E501
-
- model_config = ConfigDict(
- arbitrary_types_allowed=True,
- )
-
- def _run(
- self,
- tool_input: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- result = self.graphql_wrapper.run(tool_input)
- return json.dumps(result, indent=2)
diff --git a/libs/community/langchain_community/tools/human/__init__.py b/libs/community/langchain_community/tools/human/__init__.py
deleted file mode 100644
index 084487d0f9..0000000000
--- a/libs/community/langchain_community/tools/human/__init__.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""Tool for asking for human input."""
-
-from langchain_community.tools.human.tool import HumanInputRun
-
-__all__ = ["HumanInputRun"]
diff --git a/libs/community/langchain_community/tools/human/tool.py b/libs/community/langchain_community/tools/human/tool.py
deleted file mode 100644
index d9e238b93c..0000000000
--- a/libs/community/langchain_community/tools/human/tool.py
+++ /dev/null
@@ -1,34 +0,0 @@
-"""Tool for asking human input."""
-
-from typing import Callable, Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import Field
-
-
-def _print_func(text: str) -> None:
- print("\n") # noqa: T201
- print(text) # noqa: T201
-
-
-class HumanInputRun(BaseTool):
- """Tool that asks user for input."""
-
- name: str = "human"
- description: str = (
- "You can ask a human for guidance when you think you "
- "got stuck or you are not sure what to do next. "
- "The input should be a question for the human."
- )
- prompt_func: Callable[[str], None] = Field(default_factory=lambda: _print_func)
- input_func: Callable = Field(default_factory=lambda: input)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Human input tool."""
- self.prompt_func(query)
- return self.input_func()
diff --git a/libs/community/langchain_community/tools/ifttt.py b/libs/community/langchain_community/tools/ifttt.py
deleted file mode 100644
index 40bbe76fda..0000000000
--- a/libs/community/langchain_community/tools/ifttt.py
+++ /dev/null
@@ -1,61 +0,0 @@
-"""From https://github.com/SidU/teams-langchain-js/wiki/Connecting-IFTTT-Services.
-
-# Creating a webhook
-- Go to https://ifttt.com/create
-
-# Configuring the "If This"
-- Click on the "If This" button in the IFTTT interface.
-- Search for "Webhooks" in the search bar.
-- Choose the first option for "Receive a web request with a JSON payload."
-- Choose an Event Name that is specific to the service you plan to connect to.
-This will make it easier for you to manage the webhook URL.
-For example, if you're connecting to Spotify, you could use "Spotify" as your
-Event Name.
-- Click the "Create Trigger" button to save your settings and create your webhook.
-
-# Configuring the "Then That"
-- Tap on the "Then That" button in the IFTTT interface.
-- Search for the service you want to connect, such as Spotify.
-- Choose an action from the service, such as "Add track to a playlist".
-- Configure the action by specifying the necessary details, such as the playlist name,
-e.g., "Songs from AI".
-- Reference the JSON Payload received by the Webhook in your action. For the Spotify
-scenario, choose "{{JsonPayload}}" as your search query.
-- Tap the "Create Action" button to save your action settings.
-- Once you have finished configuring your action, click the "Finish" button to
-complete the setup.
-- Congratulations! You have successfully connected the Webhook to the desired
-service, and you're ready to start receiving data and triggering actions 🎉
-
-# Finishing up
-- To get your webhook URL go to https://ifttt.com/maker_webhooks/settings
-- Copy the IFTTT key value from there. The URL is of the form
-https://maker.ifttt.com/use/YOUR_IFTTT_KEY. Grab the YOUR_IFTTT_KEY value.
-"""
-
-from typing import Optional
-
-import requests
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-
-
-class IFTTTWebhook(BaseTool):
- """IFTTT Webhook.
-
- Args:
- name: name of the tool
- description: description of the tool
- url: url to hit with the json event.
- """
-
- url: str
-
- def _run(
- self,
- tool_input: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- body = {"this": tool_input}
- response = requests.post(self.url, data=body)
- return response.text
diff --git a/libs/community/langchain_community/tools/interaction/__init__.py b/libs/community/langchain_community/tools/interaction/__init__.py
deleted file mode 100644
index be3393362d..0000000000
--- a/libs/community/langchain_community/tools/interaction/__init__.py
+++ /dev/null
@@ -1 +0,0 @@
-"""Tools for interacting with the user."""
diff --git a/libs/community/langchain_community/tools/interaction/tool.py b/libs/community/langchain_community/tools/interaction/tool.py
deleted file mode 100644
index 6f5e84884c..0000000000
--- a/libs/community/langchain_community/tools/interaction/tool.py
+++ /dev/null
@@ -1,16 +0,0 @@
-"""Tools for interacting with the user."""
-
-import warnings
-from typing import Any
-
-from langchain_community.tools.human.tool import HumanInputRun
-
-
-def StdInInquireTool(*args: Any, **kwargs: Any) -> HumanInputRun:
- """Tool for asking the user for input."""
- warnings.warn(
- "StdInInquireTool will be deprecated in the future. "
- "Please use HumanInputRun instead.",
- DeprecationWarning,
- )
- return HumanInputRun(*args, **kwargs)
diff --git a/libs/community/langchain_community/tools/jina_search/__init__.py b/libs/community/langchain_community/tools/jina_search/__init__.py
deleted file mode 100644
index 7f924e6054..0000000000
--- a/libs/community/langchain_community/tools/jina_search/__init__.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""Jina AI toolkit"""
-
-from langchain_community.tools.jina_search.tool import JinaSearch
-
-__all__ = ["JinaSearch"]
diff --git a/libs/community/langchain_community/tools/jina_search/tool.py b/libs/community/langchain_community/tools/jina_search/tool.py
deleted file mode 100644
index 4c4e7650b6..0000000000
--- a/libs/community/langchain_community/tools/jina_search/tool.py
+++ /dev/null
@@ -1,41 +0,0 @@
-from __future__ import annotations
-
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.utilities.jina_search import JinaSearchAPIWrapper
-
-
-class JinaInput(BaseModel):
- """Input for the Jina search tool."""
-
- query: str = Field(description="search query to look up")
-
-
-class JinaSearch(BaseTool):
- """Tool that queries the JinaSearch.
-
- ..versionadded:: 0.2.16
- """
-
- name: str = "jina_search"
- description: str = (
- "Jina Reader allows you to ground your LLM with the latest information from "
- "the web. "
- "Jina Reader will search the web and return the top five results with their "
- "URLs and contents, "
- "each in clean, LLM-friendly text. This way, you can always keep your LLM "
- "up-to-date, improve its factuality, and reduce hallucinations."
- )
- search_wrapper: JinaSearchAPIWrapper = Field(default_factory=JinaSearchAPIWrapper) # type: ignore[arg-type]
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return self.search_wrapper.run(query)
diff --git a/libs/community/langchain_community/tools/jira/__init__.py b/libs/community/langchain_community/tools/jira/__init__.py
deleted file mode 100644
index 06cd8cbcd9..0000000000
--- a/libs/community/langchain_community/tools/jira/__init__.py
+++ /dev/null
@@ -1 +0,0 @@
-"""Jira Tool."""
diff --git a/libs/community/langchain_community/tools/jira/prompt.py b/libs/community/langchain_community/tools/jira/prompt.py
deleted file mode 100644
index 4e47048aa3..0000000000
--- a/libs/community/langchain_community/tools/jira/prompt.py
+++ /dev/null
@@ -1,42 +0,0 @@
-# flake8: noqa
-JIRA_ISSUE_CREATE_PROMPT = """
- This tool is a wrapper around atlassian-python-api's Jira issue_create API, useful when you need to create a Jira issue.
- The input to this tool is a dictionary specifying the fields of the Jira issue, and will be passed into atlassian-python-api's Jira `issue_create` function.
- For example, to create a low priority task called "test issue" with description "test description", you would pass in the following dictionary:
- {{"summary": "test issue", "description": "test description", "issuetype": {{"name": "Task"}}, "priority": {{"name": "Low"}}}}
- """
-
-JIRA_GET_ALL_PROJECTS_PROMPT = """
- This tool is a wrapper around atlassian-python-api's Jira project API,
- useful when you need to fetch all the projects the user has access to, find out how many projects there are, or as an intermediary step that involve searching by projects.
- there is no input to this tool.
- """
-
-JIRA_JQL_PROMPT = """
- This tool is a wrapper around atlassian-python-api's Jira jql API, useful when you need to search for Jira issues.
- The input to this tool is a JQL query string, and will be passed into atlassian-python-api's Jira `jql` function,
- For example, to find all the issues in project "Test" assigned to the me, you would pass in the following string:
- project = Test AND assignee = currentUser()
- or to find issues with summaries that contain the word "test", you would pass in the following string:
- summary ~ 'test'
- """
-
-JIRA_CATCH_ALL_PROMPT = """
- This tool is a wrapper around atlassian-python-api's Jira API.
- There are other dedicated tools for fetching all projects, and creating and searching for issues,
- use this tool if you need to perform any other actions allowed by the atlassian-python-api Jira API.
- The input to this tool is a dictionary specifying a function from atlassian-python-api's Jira API,
- as well as a list of arguments and dictionary of keyword arguments to pass into the function.
- For example, to get all the users in a group, while increasing the max number of results to 100, you would
- pass in the following dictionary: {{"function": "get_all_users_from_group", "args": ["group"], "kwargs": {{"limit":100}} }}
- or to find out how many projects are in the Jira instance, you would pass in the following string:
- {{"function": "projects"}}
- For more information on the Jira API, refer to https://atlassian-python-api.readthedocs.io/jira.html
- """
-
-JIRA_CONFLUENCE_PAGE_CREATE_PROMPT = """This tool is a wrapper around atlassian-python-api's Confluence
-atlassian-python-api API, useful when you need to create a Confluence page. The input to this tool is a dictionary
-specifying the fields of the Confluence page, and will be passed into atlassian-python-api's Confluence `create_page`
-function. For example, to create a page in the DEMO space titled "This is the title" with body "This is the body. You can use
-HTML tags!", you would pass in the following dictionary: {{"space": "DEMO", "title":"This is the
-title","body":"This is the body. You can use HTML tags!"}} """
diff --git a/libs/community/langchain_community/tools/jira/tool.py b/libs/community/langchain_community/tools/jira/tool.py
deleted file mode 100644
index 93205920c5..0000000000
--- a/libs/community/langchain_community/tools/jira/tool.py
+++ /dev/null
@@ -1,46 +0,0 @@
-"""
-This tool allows agents to interact with the atlassian-python-api library
-and operate on a Jira instance. For more information on the
-atlassian-python-api library, see https://atlassian-python-api.readthedocs.io/jira.html
-
-To use this tool, you must first set as environment variables:
- JIRA_API_TOKEN
- JIRA_USERNAME
- JIRA_INSTANCE_URL
- JIRA_CLOUD
-
-Below is a sample script that uses the Jira tool:
-
-```python
-from langchain_community.agent_toolkits.jira.toolkit import JiraToolkit
-from langchain_community.utilities.jira import JiraAPIWrapper
-
-jira = JiraAPIWrapper()
-toolkit = JiraToolkit.from_jira_api_wrapper(jira)
-```
-"""
-
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import Field
-
-from langchain_community.utilities.jira import JiraAPIWrapper
-
-
-class JiraAction(BaseTool):
- """Tool that queries the Atlassian Jira API."""
-
- api_wrapper: JiraAPIWrapper = Field(default_factory=JiraAPIWrapper)
- mode: str
- name: str = ""
- description: str = ""
-
- def _run(
- self,
- instructions: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Atlassian Jira API to run an operation."""
- return self.api_wrapper.run(self.mode, instructions)
diff --git a/libs/community/langchain_community/tools/json/__init__.py b/libs/community/langchain_community/tools/json/__init__.py
deleted file mode 100644
index d13302f008..0000000000
--- a/libs/community/langchain_community/tools/json/__init__.py
+++ /dev/null
@@ -1 +0,0 @@
-"""Tools for interacting with a JSON file."""
diff --git a/libs/community/langchain_community/tools/json/tool.py b/libs/community/langchain_community/tools/json/tool.py
deleted file mode 100644
index 6e7fddff6d..0000000000
--- a/libs/community/langchain_community/tools/json/tool.py
+++ /dev/null
@@ -1,134 +0,0 @@
-# flake8: noqa
-"""Tools for working with JSON specs."""
-
-from __future__ import annotations
-
-import json
-import re
-from pathlib import Path
-from typing import Dict, List, Optional, Union
-
-from pydantic import BaseModel
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from langchain_core.tools import BaseTool
-
-
-def _parse_input(text: str) -> List[Union[str, int]]:
- """Parse input of the form data["key1"][0]["key2"] into a list of keys."""
- _res = re.findall(r"\[.*?]", text)
- # strip the brackets and quotes, convert to int if possible
- res = [i[1:-1].replace('"', "").replace("'", "") for i in _res]
- res = [int(i) if i.isdigit() else i for i in res]
- return res
-
-
-class JsonSpec(BaseModel):
- """Base class for JSON spec."""
-
- dict_: Dict
- max_value_length: int = 200
-
- @classmethod
- def from_file(cls, path: Path) -> JsonSpec:
- """Create a JsonSpec from a file."""
- if not path.exists():
- raise FileNotFoundError(f"File not found: {path}")
- dict_ = json.loads(path.read_text())
- return cls(dict_=dict_)
-
- def keys(self, text: str) -> str:
- """Return the keys of the dict at the given path.
-
- Args:
- text: Python representation of the path to the dict (e.g. data["key1"][0]["key2"]).
- """
- try:
- items = _parse_input(text)
- val = self.dict_
- for i in items:
- if i:
- val = val[i]
- if not isinstance(val, dict):
- raise ValueError(
- f"Value at path `{text}` is not a dict, get the value directly."
- )
- return str(list(val.keys()))
- except Exception as e:
- return repr(e)
-
- def value(self, text: str) -> str:
- """Return the value of the dict at the given path.
-
- Args:
- text: Python representation of the path to the dict (e.g. data["key1"][0]["key2"]).
- """
- try:
- items = _parse_input(text)
- val = self.dict_
- for i in items:
- val = val[i]
-
- if isinstance(val, dict) and len(str(val)) > self.max_value_length:
- return "Value is a large dictionary, should explore its keys directly"
- str_val = str(val)
- if len(str_val) > self.max_value_length:
- str_val = str_val[: self.max_value_length] + "..."
- return str_val
- except Exception as e:
- return repr(e)
-
-
-class JsonListKeysTool(BaseTool):
- """Tool for listing keys in a JSON spec."""
-
- name: str = "json_spec_list_keys"
- description: str = """
- Can be used to list all keys at a given path.
- Before calling this you should be SURE that the path to this exists.
- The input is a text representation of the path to the dict in Python syntax (e.g. data["key1"][0]["key2"]).
- """
- spec: JsonSpec
-
- def _run(
- self,
- tool_input: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- return self.spec.keys(tool_input)
-
- async def _arun(
- self,
- tool_input: str,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- return self._run(tool_input)
-
-
-class JsonGetValueTool(BaseTool):
- """Tool for getting a value in a JSON spec."""
-
- name: str = "json_spec_get_value"
- description: str = """
- Can be used to see value in string format at a given path.
- Before calling this you should be SURE that the path to this exists.
- The input is a text representation of the path to the dict in Python syntax (e.g. data["key1"][0]["key2"]).
- """
- spec: JsonSpec
-
- def _run(
- self,
- tool_input: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- return self.spec.value(tool_input)
-
- async def _arun(
- self,
- tool_input: str,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- return self._run(tool_input)
diff --git a/libs/community/langchain_community/tools/memorize/__init__.py b/libs/community/langchain_community/tools/memorize/__init__.py
deleted file mode 100644
index 76a84406ac..0000000000
--- a/libs/community/langchain_community/tools/memorize/__init__.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""Unsupervised learning based memorization."""
-
-from langchain_community.tools.memorize.tool import Memorize
-
-__all__ = ["Memorize"]
diff --git a/libs/community/langchain_community/tools/memorize/tool.py b/libs/community/langchain_community/tools/memorize/tool.py
deleted file mode 100644
index 87badf9ac3..0000000000
--- a/libs/community/langchain_community/tools/memorize/tool.py
+++ /dev/null
@@ -1,60 +0,0 @@
-from abc import abstractmethod
-from typing import Any, Optional, Protocol, Sequence, runtime_checkable
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from langchain_core.tools import BaseTool
-from pydantic import Field
-
-from langchain_community.llms.gradient_ai import TrainResult
-
-
-@runtime_checkable
-class TrainableLLM(Protocol):
- """Protocol for trainable language models."""
-
- @abstractmethod
- def train_unsupervised(
- self,
- inputs: Sequence[str],
- **kwargs: Any,
- ) -> TrainResult: ...
-
- @abstractmethod
- async def atrain_unsupervised(
- self,
- inputs: Sequence[str],
- **kwargs: Any,
- ) -> TrainResult: ...
-
-
-class Memorize(BaseTool):
- """Tool that trains a language model."""
-
- name: str = "memorize"
- description: str = (
- "Useful whenever you observed novel information "
- "from previous conversation history, "
- "i.e., another tool's action outputs or human comments. "
- "The action input should include observed information in detail, "
- "then the tool will fine-tune yourself to remember it."
- )
- llm: TrainableLLM = Field()
-
- def _run(
- self,
- information_to_learn: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- train_result = self.llm.train_unsupervised((information_to_learn,))
- return f"Train complete. Loss: {train_result['loss']}"
-
- async def _arun(
- self,
- information_to_learn: str,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- train_result = await self.llm.atrain_unsupervised((information_to_learn,))
- return f"Train complete. Loss: {train_result['loss']}"
diff --git a/libs/community/langchain_community/tools/merriam_webster/__init__.py b/libs/community/langchain_community/tools/merriam_webster/__init__.py
deleted file mode 100644
index 73390d5498..0000000000
--- a/libs/community/langchain_community/tools/merriam_webster/__init__.py
+++ /dev/null
@@ -1 +0,0 @@
-"""Merriam-Webster API toolkit."""
diff --git a/libs/community/langchain_community/tools/merriam_webster/tool.py b/libs/community/langchain_community/tools/merriam_webster/tool.py
deleted file mode 100644
index 9cf4e9f21c..0000000000
--- a/libs/community/langchain_community/tools/merriam_webster/tool.py
+++ /dev/null
@@ -1,28 +0,0 @@
-"""Tool for the Merriam-Webster API."""
-
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-
-from langchain_community.utilities.merriam_webster import MerriamWebsterAPIWrapper
-
-
-class MerriamWebsterQueryRun(BaseTool):
- """Tool that searches the Merriam-Webster API."""
-
- name: str = "merriam_webster"
- description: str = (
- "A wrapper around Merriam-Webster. "
- "Useful for when you need to get the definition of a word."
- "Input should be the word you want the definition of."
- )
- api_wrapper: MerriamWebsterAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Merriam-Webster tool."""
- return self.api_wrapper.run(query)
diff --git a/libs/community/langchain_community/tools/metaphor_search/__init__.py b/libs/community/langchain_community/tools/metaphor_search/__init__.py
deleted file mode 100644
index 246f25a129..0000000000
--- a/libs/community/langchain_community/tools/metaphor_search/__init__.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""Metaphor Search API toolkit."""
-
-from langchain_community.tools.metaphor_search.tool import MetaphorSearchResults
-
-__all__ = ["MetaphorSearchResults"]
diff --git a/libs/community/langchain_community/tools/metaphor_search/tool.py b/libs/community/langchain_community/tools/metaphor_search/tool.py
deleted file mode 100644
index 98e932e8d2..0000000000
--- a/libs/community/langchain_community/tools/metaphor_search/tool.py
+++ /dev/null
@@ -1,87 +0,0 @@
-"""Tool for the Metaphor search API."""
-
-from typing import Dict, List, Optional, Union
-
-from langchain_core._api.deprecation import deprecated
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from langchain_core.tools import BaseTool
-
-from langchain_community.utilities.metaphor_search import MetaphorSearchAPIWrapper
-
-
-@deprecated(
- since="0.0.15",
- removal="1.0",
- alternative="langchain_exa.ExaSearchResults",
-)
-class MetaphorSearchResults(BaseTool):
- """Tool that queries the Metaphor Search API and gets back json."""
-
- name: str = "metaphor_search_results_json"
- description: str = (
- "A wrapper around Metaphor Search. "
- "Input should be a Metaphor-optimized query. "
- "Output is a JSON array of the query results"
- )
- api_wrapper: MetaphorSearchAPIWrapper
-
- def _run(
- self,
- query: str,
- num_results: int,
- include_domains: Optional[List[str]] = None,
- exclude_domains: Optional[List[str]] = None,
- start_crawl_date: Optional[str] = None,
- end_crawl_date: Optional[str] = None,
- start_published_date: Optional[str] = None,
- end_published_date: Optional[str] = None,
- use_autoprompt: Optional[bool] = None,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> Union[List[Dict], str]:
- """Use the tool."""
- try:
- return self.api_wrapper.results(
- query,
- num_results,
- include_domains,
- exclude_domains,
- start_crawl_date,
- end_crawl_date,
- start_published_date,
- end_published_date,
- use_autoprompt,
- )
- except Exception as e:
- return repr(e)
-
- async def _arun(
- self,
- query: str,
- num_results: int,
- include_domains: Optional[List[str]] = None,
- exclude_domains: Optional[List[str]] = None,
- start_crawl_date: Optional[str] = None,
- end_crawl_date: Optional[str] = None,
- start_published_date: Optional[str] = None,
- end_published_date: Optional[str] = None,
- use_autoprompt: Optional[bool] = None,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> Union[List[Dict], str]:
- """Use the tool asynchronously."""
- try:
- return await self.api_wrapper.results_async(
- query,
- num_results,
- include_domains,
- exclude_domains,
- start_crawl_date,
- end_crawl_date,
- start_published_date,
- end_published_date,
- use_autoprompt,
- )
- except Exception as e:
- return repr(e)
diff --git a/libs/community/langchain_community/tools/mojeek_search/__init__.py b/libs/community/langchain_community/tools/mojeek_search/__init__.py
deleted file mode 100644
index e69de29bb2..0000000000
diff --git a/libs/community/langchain_community/tools/mojeek_search/tool.py b/libs/community/langchain_community/tools/mojeek_search/tool.py
deleted file mode 100644
index 9112e1afe6..0000000000
--- a/libs/community/langchain_community/tools/mojeek_search/tool.py
+++ /dev/null
@@ -1,45 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Optional
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from langchain_core.tools import BaseTool
-
-from langchain_community.utilities.mojeek_search import MojeekSearchAPIWrapper
-
-
-class MojeekSearch(BaseTool):
- name: str = "mojeek_search"
- description: str = (
- "A wrapper around Mojeek Search. "
- "Useful for when you need to web search results. "
- "Input should be a search query."
- )
- api_wrapper: MojeekSearchAPIWrapper
-
- @classmethod
- def config(
- cls, api_key: str, search_kwargs: Optional[dict] = None, **kwargs: Any
- ) -> MojeekSearch:
- wrapper = MojeekSearchAPIWrapper(
- api_key=api_key, search_kwargs=search_kwargs or {}
- )
- return cls(api_wrapper=wrapper, **kwargs)
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- return self.api_wrapper.run(query)
-
- async def _arun(
- self,
- query: str,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool asynchronously."""
- raise NotImplementedError("MojeekSearch does not support async")
diff --git a/libs/community/langchain_community/tools/multion/__init__.py b/libs/community/langchain_community/tools/multion/__init__.py
deleted file mode 100644
index c273a08861..0000000000
--- a/libs/community/langchain_community/tools/multion/__init__.py
+++ /dev/null
@@ -1,7 +0,0 @@
-"""MutliOn Client API tools."""
-
-from langchain_community.tools.multion.close_session import MultionCloseSession
-from langchain_community.tools.multion.create_session import MultionCreateSession
-from langchain_community.tools.multion.update_session import MultionUpdateSession
-
-__all__ = ["MultionCreateSession", "MultionUpdateSession", "MultionCloseSession"]
diff --git a/libs/community/langchain_community/tools/multion/close_session.py b/libs/community/langchain_community/tools/multion/close_session.py
deleted file mode 100644
index 28f0abd013..0000000000
--- a/libs/community/langchain_community/tools/multion/close_session.py
+++ /dev/null
@@ -1,57 +0,0 @@
-from typing import TYPE_CHECKING, Optional, Type
-
-from langchain_core.callbacks import (
- CallbackManagerForToolRun,
-)
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-if TYPE_CHECKING:
- # This is for linting and IDE typehints
- import multion
-else:
- try:
- # We do this so pydantic can resolve the types when instantiating
- import multion
- except ImportError:
- pass
-
-
-class CloseSessionSchema(BaseModel):
- """Input for UpdateSessionTool."""
-
- sessionId: str = Field(
- ...,
- description="""The sessionId, received from one of the createSessions
- or updateSessions run before""",
- )
-
-
-class MultionCloseSession(BaseTool):
- """Tool that closes an existing Multion Browser Window with provided fields.
-
- Attributes:
- name: The name of the tool. Default: "close_multion_session"
- description: The description of the tool.
- args_schema: The schema for the tool's arguments. Default: UpdateSessionSchema
- """
-
- name: str = "close_multion_session"
- description: str = """Use this tool to close \
-an existing corresponding Multion Browser Window with provided fields. \
-Note: SessionId must be received from previous Browser window creation."""
- args_schema: Type[CloseSessionSchema] = CloseSessionSchema
- sessionId: str = ""
-
- def _run(
- self,
- sessionId: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> None:
- try:
- try:
- multion.close_session(sessionId)
- except Exception as e:
- print(f"{e}, retrying...") # noqa: T201
- except Exception as e:
- raise Exception(f"An error occurred: {e}")
diff --git a/libs/community/langchain_community/tools/multion/create_session.py b/libs/community/langchain_community/tools/multion/create_session.py
deleted file mode 100644
index 53388a5a97..0000000000
--- a/libs/community/langchain_community/tools/multion/create_session.py
+++ /dev/null
@@ -1,67 +0,0 @@
-from typing import TYPE_CHECKING, Optional, Type
-
-from langchain_core.callbacks import (
- CallbackManagerForToolRun,
-)
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-if TYPE_CHECKING:
- # This is for linting and IDE typehints
- import multion
-else:
- try:
- # We do this so pydantic can resolve the types when instantiating
- import multion
- except ImportError:
- pass
-
-
-class CreateSessionSchema(BaseModel):
- """Input for CreateSessionTool."""
-
- query: str = Field(
- ...,
- description="The query to run in multion agent.",
- )
- url: str = Field(
- "https://www.google.com/",
- description="""The Url to run the agent at. Note: accepts only secure \
- links having https://""",
- )
-
-
-class MultionCreateSession(BaseTool):
- """Tool that creates a new Multion Browser Window with provided fields.
-
- Attributes:
- name: The name of the tool. Default: "create_multion_session"
- description: The description of the tool.
- args_schema: The schema for the tool's arguments.
- """
-
- name: str = "create_multion_session"
- description: str = """
- Create a new web browsing session based on a user's command or request. \
- The command should include the full info required for the session. \
- Also include an url (defaults to google.com if no better option) \
- to start the session. \
- Use this tool to create a new Browser Window with provided fields. \
- Always the first step to run any activities that can be done using browser.
- """
- args_schema: Type[CreateSessionSchema] = CreateSessionSchema
-
- def _run(
- self,
- query: str,
- url: Optional[str] = "https://www.google.com/",
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> dict:
- try:
- response = multion.new_session({"input": query, "url": url})
- return {
- "sessionId": response["session_id"],
- "Response": response["message"],
- }
- except Exception as e:
- raise Exception(f"An error occurred: {e}")
diff --git a/libs/community/langchain_community/tools/multion/update_session.py b/libs/community/langchain_community/tools/multion/update_session.py
deleted file mode 100644
index b535861e2a..0000000000
--- a/libs/community/langchain_community/tools/multion/update_session.py
+++ /dev/null
@@ -1,74 +0,0 @@
-from typing import TYPE_CHECKING, Optional, Type
-
-from langchain_core.callbacks import (
- CallbackManagerForToolRun,
-)
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-if TYPE_CHECKING:
- # This is for linting and IDE typehints
- import multion
-else:
- try:
- # We do this so pydantic can resolve the types when instantiating
- import multion
- except ImportError:
- pass
-
-
-class UpdateSessionSchema(BaseModel):
- """Input for UpdateSessionTool."""
-
- sessionId: str = Field(
- ...,
- description="""The sessionID,
- received from one of the createSessions run before""",
- )
- query: str = Field(
- ...,
- description="The query to run in multion agent.",
- )
- url: str = Field(
- "https://www.google.com/",
- description="""The Url to run the agent at. \
- Note: accepts only secure links having https://""",
- )
-
-
-class MultionUpdateSession(BaseTool):
- """Tool that updates an existing Multion Browser Window with provided fields.
-
- Attributes:
- name: The name of the tool. Default: "update_multion_session"
- description: The description of the tool.
- args_schema: The schema for the tool's arguments. Default: UpdateSessionSchema
- """
-
- name: str = "update_multion_session"
- description: str = """Use this tool to update \
-an existing corresponding Multion Browser Window with provided fields. \
-Note: sessionId must be received from previous Browser window creation."""
- args_schema: Type[UpdateSessionSchema] = UpdateSessionSchema
- sessionId: str = ""
-
- def _run(
- self,
- sessionId: str,
- query: str,
- url: Optional[str] = "https://www.google.com/",
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> dict:
- try:
- try:
- response = multion.update_session(
- sessionId, {"input": query, "url": url}
- )
- content = {"sessionId": sessionId, "Response": response["message"]}
- self.sessionId = sessionId
- return content
- except Exception as e:
- print(f"{e}, retrying...") # noqa: T201
- return {"error": f"{e}", "Response": "retrying..."}
- except Exception as e:
- raise Exception(f"An error occurred: {e}")
diff --git a/libs/community/langchain_community/tools/nasa/__init__.py b/libs/community/langchain_community/tools/nasa/__init__.py
deleted file mode 100644
index e69de29bb2..0000000000
diff --git a/libs/community/langchain_community/tools/nasa/prompt.py b/libs/community/langchain_community/tools/nasa/prompt.py
deleted file mode 100644
index 4c7a3846a7..0000000000
--- a/libs/community/langchain_community/tools/nasa/prompt.py
+++ /dev/null
@@ -1,82 +0,0 @@
-# flake8: noqa
-NASA_SEARCH_PROMPT = """
- This tool is a wrapper around NASA's search API, useful when you need to search through NASA's Image and Video Library.
- The input to this tool is a query specified by the user, and will be passed into NASA's `search` function.
-
- At least one parameter must be provided.
-
- There are optional parameters that can be passed by the user based on their query
- specifications. Each item in this list contains pound sign (#) separated values, the first value is the parameter name,
- the second value is the datatype and the third value is the description: {{
-
- - q#string#Free text search terms to compare to all indexed metadata.
- - center#string#NASA center which published the media.
- - description#string#Terms to search for in “Description” fields.
- - description_508#string#Terms to search for in “508 Description” fields.
- - keywords #string#Terms to search for in “Keywords” fields. Separate multiple values with commas.
- - location #string#Terms to search for in “Location” fields.
- - media_type#string#Media types to restrict the search to. Available types: [“image”,“video”, “audio”]. Separate multiple values with commas.
- - nasa_id #string#The media asset’s NASA ID.
- - page#integer#Page number, starting at 1, of results to get.-
- - page_size#integer#Number of results per page. Default: 100.
- - photographer#string#The primary photographer’s name.
- - secondary_creator#string#A secondary photographer/videographer’s name.
- - title #string#Terms to search for in “Title” fields.
- - year_start#string#The start year for results. Format: YYYY.
- - year_end #string#The end year for results. Format: YYYY.
-
- }}
-
- Below are several task descriptions along with their respective input examples.
- Task: get the 2nd page of image and video content starting from the year 2002 to 2010
- Example Input: {{"year_start": "2002", "year_end": "2010", "page": 2}}
-
- Task: get the image and video content of saturn photographed by John Appleseed
- Example Input: {{"q": "saturn", "photographer": "John Appleseed"}}
-
- Task: search for Meteor Showers with description "Search Description" with media type image
- Example Input: {{"q": "Meteor Shower", "description": "Search Description", "media_type": "image"}}
-
- Task: get the image and video content from year 2008 to 2010 from Kennedy Center
- Example Input: {{"year_start": "2002", "year_end": "2010", "location": "Kennedy Center}}
- """
-
-
-NASA_MANIFEST_PROMPT = """
- This tool is a wrapper around NASA's media asset manifest API, useful when you need to retrieve a media
- asset's manifest. The input to this tool should include a string representing a NASA ID for a media asset that the user is trying to get the media asset manifest data for. The NASA ID will be passed as a string into NASA's `get_media_metadata_manifest` function.
-
- The following list are some examples of NASA IDs for a media asset that you can use to better extract the NASA ID from the input string to the tool.
- - GSFC_20171102_Archive_e000579
- - Launch-Sound_Delta-PAM-Random-Commentary
- - iss066m260341519_Expedition_66_Education_Inflight_with_Random_Lake_School_District_220203
- - 6973610
- - GRC-2020-CM-0167.4
- - Expedition_55_Inflight_Japan_VIP_Event_May_31_2018_659970
- - NASA 60th_SEAL_SLIVER_150DPI
-"""
-
-NASA_METADATA_PROMPT = """
- This tool is a wrapper around NASA's media asset metadata location API, useful when you need to retrieve the media asset's metadata. The input to this tool should include a string representing a NASA ID for a media asset that the user is trying to get the media asset metadata location for. The NASA ID will be passed as a string into NASA's `get_media_metadata_manifest` function.
-
- The following list are some examples of NASA IDs for a media asset that you can use to better extract the NASA ID from the input string to the tool.
- - GSFC_20171102_Archive_e000579
- - Launch-Sound_Delta-PAM-Random-Commentary
- - iss066m260341519_Expedition_66_Education_Inflight_with_Random_Lake_School_District_220203
- - 6973610
- - GRC-2020-CM-0167.4
- - Expedition_55_Inflight_Japan_VIP_Event_May_31_2018_659970
- - NASA 60th_SEAL_SLIVER_150DPI
-"""
-
-NASA_CAPTIONS_PROMPT = """
- This tool is a wrapper around NASA's video assests caption location API, useful when you need
- to retrieve the location of the captions of a specific video. The input to this tool should include a string representing a NASA ID for a video media asset that the user is trying to get the get the location of the captions for. The NASA ID will be passed as a string into NASA's `get_media_metadata_manifest` function.
-
- The following list are some examples of NASA IDs for a video asset that you can use to better extract the NASA ID from the input string to the tool.
- - 2017-08-09 - Video File RS-25 Engine Test
- - 20180415-TESS_Social_Briefing
- - 201_TakingWildOutOfWildfire
- - 2022-H1_V_EuropaClipper-4
- - 2022_0429_Recientemente
-"""
diff --git a/libs/community/langchain_community/tools/nasa/tool.py b/libs/community/langchain_community/tools/nasa/tool.py
deleted file mode 100644
index b9f2caa455..0000000000
--- a/libs/community/langchain_community/tools/nasa/tool.py
+++ /dev/null
@@ -1,29 +0,0 @@
-"""
-This tool allows agents to interact with the NASA API, specifically
-the the NASA Image & Video Library and Exoplanet
-"""
-
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import Field
-
-from langchain_community.utilities.nasa import NasaAPIWrapper
-
-
-class NasaAction(BaseTool):
- """Tool that queries the Atlassian Jira API."""
-
- api_wrapper: NasaAPIWrapper = Field(default_factory=NasaAPIWrapper)
- mode: str
- name: str = ""
- description: str = ""
-
- def _run(
- self,
- instructions: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the NASA API to run an operation."""
- return self.api_wrapper.run(self.mode, instructions)
diff --git a/libs/community/langchain_community/tools/nuclia/__init__.py b/libs/community/langchain_community/tools/nuclia/__init__.py
deleted file mode 100644
index ea2f3dc651..0000000000
--- a/libs/community/langchain_community/tools/nuclia/__init__.py
+++ /dev/null
@@ -1,3 +0,0 @@
-from langchain_community.tools.nuclia.tool import NucliaUnderstandingAPI
-
-__all__ = ["NucliaUnderstandingAPI"]
diff --git a/libs/community/langchain_community/tools/nuclia/tool.py b/libs/community/langchain_community/tools/nuclia/tool.py
deleted file mode 100644
index 8aeed0feb3..0000000000
--- a/libs/community/langchain_community/tools/nuclia/tool.py
+++ /dev/null
@@ -1,237 +0,0 @@
-"""Tool for the Nuclia Understanding API.
-
-Installation:
-
-```bash
- pip install --upgrade protobuf
- pip install nucliadb-protos
-```
-"""
-
-import asyncio
-import base64
-import logging
-import mimetypes
-import os
-from typing import Any, Dict, Optional, Type, Union
-
-import requests
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-logger = logging.getLogger(__name__)
-
-
-class NUASchema(BaseModel):
- """Input for Nuclia Understanding API.
-
- Attributes:
- action: Action to perform. Either `push` or `pull`.
- id: ID of the file to push or pull.
- path: Path to the file to push (needed only for `push` action).
- text: Text content to process (needed only for `push` action).
- """
-
- action: str = Field(
- ...,
- description="Action to perform. Either `push` or `pull`.",
- )
- id: str = Field(
- ...,
- description="ID of the file to push or pull.",
- )
- path: Optional[str] = Field(
- ...,
- description="Path to the file to push (needed only for `push` action).",
- )
- text: Optional[str] = Field(
- ...,
- description="Text content to process (needed only for `push` action).",
- )
-
-
-class NucliaUnderstandingAPI(BaseTool):
- """Tool to process files with the Nuclia Understanding API."""
-
- name: str = "nuclia_understanding_api"
- description: str = (
- "A wrapper around Nuclia Understanding API endpoints. "
- "Useful for when you need to extract text from any kind of files. "
- )
- args_schema: Type[BaseModel] = NUASchema
- _results: Dict[str, Any] = {}
- _config: Dict[str, Any] = {}
-
- def __init__(self, enable_ml: bool = False) -> None:
- zone = os.environ.get("NUCLIA_ZONE", "europe-1")
- self._config["BACKEND"] = f"https://{zone}.nuclia.cloud/api/v1"
- key = os.environ.get("NUCLIA_NUA_KEY")
- if not key:
- raise ValueError("NUCLIA_NUA_KEY environment variable not set")
- else:
- self._config["NUA_KEY"] = key
- self._config["enable_ml"] = enable_ml
- super().__init__()
-
- def _run(
- self,
- action: str,
- id: str,
- path: Optional[str],
- text: Optional[str],
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- if action == "push":
- self._check_params(path, text)
- if path:
- return self._pushFile(id, path)
- if text:
- return self._pushText(id, text)
- elif action == "pull":
- return self._pull(id)
- return ""
-
- async def _arun(
- self,
- action: str,
- id: str,
- path: Optional[str] = None,
- text: Optional[str] = None,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool asynchronously."""
- self._check_params(path, text)
- if path:
- self._pushFile(id, path)
- if text:
- self._pushText(id, text)
- data = None
- while True:
- data = self._pull(id)
- if data:
- break
- await asyncio.sleep(15)
- return data
-
- def _pushText(self, id: str, text: str) -> str:
- field = {
- "textfield": {"text": {"body": text, "format": 0}},
- "processing_options": {"ml_text": self._config["enable_ml"]},
- }
- return self._pushField(id, field)
-
- def _pushFile(self, id: str, content_path: str) -> str:
- with open(content_path, "rb") as source_file:
- response = requests.post(
- self._config["BACKEND"] + "/processing/upload",
- headers={
- "content-type": mimetypes.guess_type(content_path)[0]
- or "application/octet-stream",
- "x-stf-nuakey": "Bearer " + self._config["NUA_KEY"],
- },
- data=source_file.read(),
- )
- if response.status_code != 200:
- logger.info(
- f"Error uploading {content_path}: "
- f"{response.status_code} {response.text}"
- )
- return ""
- else:
- field = {
- "filefield": {"file": f"{response.text}"},
- "processing_options": {"ml_text": self._config["enable_ml"]},
- }
- return self._pushField(id, field)
-
- def _pushField(self, id: str, field: Any) -> str:
- logger.info(f"Pushing {id} in queue")
- response = requests.post(
- self._config["BACKEND"] + "/processing/push",
- headers={
- "content-type": "application/json",
- "x-stf-nuakey": "Bearer " + self._config["NUA_KEY"],
- },
- json=field,
- )
- if response.status_code != 200:
- logger.info(
- f"Error pushing field {id}:{response.status_code} {response.text}"
- )
- raise ValueError("Error pushing field")
- else:
- uuid = response.json()["uuid"]
- logger.info(f"Field {id} pushed in queue, uuid: {uuid}")
- self._results[id] = {"uuid": uuid, "status": "pending"}
- return uuid
-
- def _pull(self, id: str) -> str:
- self._pull_queue()
- result = self._results.get(id, None)
- if not result:
- logger.info(f"{id} not in queue")
- return ""
- elif result["status"] == "pending":
- logger.info(f"Waiting for {result['uuid']} to be processed")
- return ""
- else:
- return result["data"]
-
- def _pull_queue(self) -> None:
- try:
- from nucliadb_protos.writer_pb2 import BrokerMessage
- except ImportError as e:
- raise ImportError(
- "nucliadb-protos is not installed. "
- "Run `pip install nucliadb-protos` to install."
- ) from e
- try:
- from google.protobuf.json_format import MessageToJson
- except ImportError as e:
- raise ImportError(
- "Unable to import google.protobuf, please install with "
- "`pip install protobuf`."
- ) from e
-
- res = requests.get(
- self._config["BACKEND"] + "/processing/pull",
- headers={
- "x-stf-nuakey": "Bearer " + self._config["NUA_KEY"],
- },
- ).json()
- if res["status"] == "empty":
- logger.info("Queue empty")
- elif res["status"] == "ok":
- payload = res["payload"]
- pb = BrokerMessage()
- pb.ParseFromString(base64.b64decode(payload))
- uuid = pb.uuid
- logger.info(f"Pulled {uuid} from queue")
- matching_id = self._find_matching_id(uuid)
- if not matching_id:
- logger.info(f"No matching id for {uuid}")
- else:
- self._results[matching_id]["status"] = "done"
- data = MessageToJson( # type: ignore[call-arg]
- pb,
- preserving_proto_field_name=True,
- including_default_value_fields=True,
- )
- self._results[matching_id]["data"] = data
-
- def _find_matching_id(self, uuid: str) -> Union[str, None]:
- for id, result in self._results.items():
- if result["uuid"] == uuid:
- return id
- return None
-
- def _check_params(self, path: Optional[str], text: Optional[str]) -> None:
- if not path and not text:
- raise ValueError("File path or text is required")
- if path and text:
- raise ValueError("Cannot process both file and text on a single run")
diff --git a/libs/community/langchain_community/tools/office365/__init__.py b/libs/community/langchain_community/tools/office365/__init__.py
deleted file mode 100644
index 10ff2206f7..0000000000
--- a/libs/community/langchain_community/tools/office365/__init__.py
+++ /dev/null
@@ -1,19 +0,0 @@
-"""O365 tools."""
-
-from langchain_community.tools.office365.create_draft_message import (
- O365CreateDraftMessage,
-)
-from langchain_community.tools.office365.events_search import O365SearchEvents
-from langchain_community.tools.office365.messages_search import O365SearchEmails
-from langchain_community.tools.office365.send_event import O365SendEvent
-from langchain_community.tools.office365.send_message import O365SendMessage
-from langchain_community.tools.office365.utils import authenticate
-
-__all__ = [
- "O365SearchEmails",
- "O365SearchEvents",
- "O365CreateDraftMessage",
- "O365SendMessage",
- "O365SendEvent",
- "authenticate",
-]
diff --git a/libs/community/langchain_community/tools/office365/base.py b/libs/community/langchain_community/tools/office365/base.py
deleted file mode 100644
index 55160bd5e5..0000000000
--- a/libs/community/langchain_community/tools/office365/base.py
+++ /dev/null
@@ -1,20 +0,0 @@
-"""Base class for Office 365 tools."""
-
-from __future__ import annotations
-
-from typing import TYPE_CHECKING
-
-from langchain_core.tools import BaseTool
-from pydantic import Field
-
-from langchain_community.tools.office365.utils import authenticate
-
-if TYPE_CHECKING:
- from O365 import Account
-
-
-class O365BaseTool(BaseTool):
- """Base class for the Office 365 tools."""
-
- account: Account = Field(default_factory=authenticate)
- """The account object for the Office 365 account."""
diff --git a/libs/community/langchain_community/tools/office365/create_draft_message.py b/libs/community/langchain_community/tools/office365/create_draft_message.py
deleted file mode 100644
index 02915ffedc..0000000000
--- a/libs/community/langchain_community/tools/office365/create_draft_message.py
+++ /dev/null
@@ -1,68 +0,0 @@
-from typing import List, Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.office365.base import O365BaseTool
-
-
-class CreateDraftMessageSchema(BaseModel):
- """Input for SendMessageTool."""
-
- body: str = Field(
- ...,
- description="The message body to include in the draft.",
- )
- to: List[str] = Field(
- ...,
- description="The list of recipients.",
- )
- subject: str = Field(
- ...,
- description="The subject of the message.",
- )
- cc: Optional[List[str]] = Field(
- None,
- description="The list of CC recipients.",
- )
- bcc: Optional[List[str]] = Field(
- None,
- description="The list of BCC recipients.",
- )
-
-
-class O365CreateDraftMessage(O365BaseTool):
- """Tool for creating a draft email in Office 365."""
-
- name: str = "create_email_draft"
- description: str = (
- "Use this tool to create a draft email with the provided message fields."
- )
- args_schema: Type[CreateDraftMessageSchema] = CreateDraftMessageSchema
-
- def _run(
- self,
- body: str,
- to: List[str],
- subject: str,
- cc: Optional[List[str]] = None,
- bcc: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- # Get mailbox object
- mailbox = self.account.mailbox()
- message = mailbox.new_message()
-
- # Assign message values
- message.body = body
- message.subject = subject
- message.to.add(to)
- if cc is not None:
- message.cc.add(cc)
- if bcc is not None:
- message.bcc.add(bcc)
-
- message.save_draft()
-
- output = "Draft created: " + str(message)
- return output
diff --git a/libs/community/langchain_community/tools/office365/events_search.py b/libs/community/langchain_community/tools/office365/events_search.py
deleted file mode 100644
index f23dd86b08..0000000000
--- a/libs/community/langchain_community/tools/office365/events_search.py
+++ /dev/null
@@ -1,127 +0,0 @@
-"""Util that Searches calendar events in Office 365.
-
-Free, but setup is required. See link below.
-https://learn.microsoft.com/en-us/graph/auth/
-"""
-
-from datetime import datetime as dt
-from typing import Any, Dict, List, Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, ConfigDict, Field
-
-from langchain_community.tools.office365.base import O365BaseTool
-from langchain_community.tools.office365.utils import UTC_FORMAT, clean_body
-
-
-class SearchEventsInput(BaseModel):
- """Input for SearchEmails Tool.
-
- From https://learn.microsoft.com/en-us/graph/search-query-parameter"""
-
- start_datetime: str = Field(
- description=(
- " The start datetime for the search query in the following format: "
- ' YYYY-MM-DDTHH:MM:SS±hh:mm, where "T" separates the date and time '
- " components, and the time zone offset is specified as ±hh:mm. "
- ' For example: "2023-06-09T10:30:00+03:00" represents June 9th, '
- " 2023, at 10:30 AM in a time zone with a positive offset of 3 "
- " hours from Coordinated Universal Time (UTC)."
- )
- )
- end_datetime: str = Field(
- description=(
- " The end datetime for the search query in the following format: "
- ' YYYY-MM-DDTHH:MM:SS±hh:mm, where "T" separates the date and time '
- " components, and the time zone offset is specified as ±hh:mm. "
- ' For example: "2023-06-09T10:30:00+03:00" represents June 9th, '
- " 2023, at 10:30 AM in a time zone with a positive offset of 3 "
- " hours from Coordinated Universal Time (UTC)."
- )
- )
- max_results: int = Field(
- default=10,
- description="The maximum number of results to return.",
- )
- truncate: bool = Field(
- default=True,
- description=(
- "Whether the event's body is truncated to meet token number limits. Set to "
- "False for searches that will retrieve small events, otherwise, set to "
- "True."
- ),
- )
-
-
-class O365SearchEvents(O365BaseTool):
- """Search calendar events in Office 365.
-
- Free, but setup is required
- """
-
- name: str = "events_search"
- args_schema: Type[BaseModel] = SearchEventsInput
- description: str = (
- " Use this tool to search for the user's calendar events."
- " The input must be the start and end datetimes for the search query."
- " The output is a JSON list of all the events in the user's calendar"
- " between the start and end times. You can assume that the user can "
- " not schedule any meeting over existing meetings, and that the user "
- "is busy during meetings. Any times without events are free for the user. "
- )
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- def _run(
- self,
- start_datetime: str,
- end_datetime: str,
- max_results: int = 10,
- truncate: bool = True,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- truncate_limit: int = 150,
- ) -> List[Dict[str, Any]]:
- # Get calendar object
- schedule = self.account.schedule()
- calendar = schedule.get_default_calendar()
-
- # Process the date range parameters
- start_datetime_query = dt.strptime(start_datetime, UTC_FORMAT)
- end_datetime_query = dt.strptime(end_datetime, UTC_FORMAT)
-
- # Run the query
- q = calendar.new_query("start").greater_equal(start_datetime_query)
- q.chain("and").on_attribute("end").less_equal(end_datetime_query)
- events = calendar.get_events(query=q, include_recurring=True, limit=max_results)
-
- # Generate output dict
- output_events = []
- for event in events:
- output_event = {}
- output_event["organizer"] = event.organizer
-
- output_event["subject"] = event.subject
-
- if truncate:
- output_event["body"] = clean_body(event.body)[:truncate_limit]
- else:
- output_event["body"] = clean_body(event.body)
-
- # Get the time zone from the search parameters
- time_zone = start_datetime_query.tzinfo
- # Assign the datetimes in the search time zone
- output_event["start_datetime"] = event.start.astimezone(time_zone).strftime(
- UTC_FORMAT
- )
- output_event["end_datetime"] = event.end.astimezone(time_zone).strftime(
- UTC_FORMAT
- )
- output_event["modified_date"] = event.modified.astimezone(
- time_zone
- ).strftime(UTC_FORMAT)
-
- output_events.append(output_event)
-
- return output_events
diff --git a/libs/community/langchain_community/tools/office365/messages_search.py b/libs/community/langchain_community/tools/office365/messages_search.py
deleted file mode 100644
index 71fe2562bb..0000000000
--- a/libs/community/langchain_community/tools/office365/messages_search.py
+++ /dev/null
@@ -1,122 +0,0 @@
-"""Util that Searches email messages in Office 365.
-
-Free, but setup is required. See link below.
-https://learn.microsoft.com/en-us/graph/auth/
-"""
-
-from typing import Any, Dict, List, Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, ConfigDict, Field
-
-from langchain_community.tools.office365.base import O365BaseTool
-from langchain_community.tools.office365.utils import UTC_FORMAT, clean_body
-
-
-class SearchEmailsInput(BaseModel):
- """Input for SearchEmails Tool."""
-
- """From https://learn.microsoft.com/en-us/graph/search-query-parameter"""
-
- folder: str = Field(
- default="",
- description=(
- " If the user wants to search in only one folder, the name of the folder. "
- 'Default folders are "inbox", "drafts", "sent items", "deleted ttems", but '
- "users can search custom folders as well."
- ),
- )
- query: str = Field(
- description=(
- "The Microsoift Graph v1.0 $search query. Example filters include "
- "from:sender, from:sender, to:recipient, subject:subject, "
- "recipients:list_of_recipients, body:excitement, importance:high, "
- "received>2022-12-01, received<2021-12-01, sent>2022-12-01, "
- "sent<2021-12-01, hasAttachments:true attachment:api-catalog.md, "
- "cc:samanthab@contoso.com, bcc:samanthab@contoso.com, body:excitement date "
- "range example: received:2023-06-08..2023-06-09 matching example: "
- "from:amy OR from:david."
- )
- )
- max_results: int = Field(
- default=10,
- description="The maximum number of results to return.",
- )
- truncate: bool = Field(
- default=True,
- description=(
- "Whether the email body is truncated to meet token number limits. Set to "
- "False for searches that will retrieve small messages, otherwise, set to "
- "True"
- ),
- )
-
-
-class O365SearchEmails(O365BaseTool):
- """Search email messages in Office 365.
-
- Free, but setup is required.
- """
-
- name: str = "messages_search"
- args_schema: Type[BaseModel] = SearchEmailsInput
- description: str = (
- "Use this tool to search for email messages."
- " The input must be a valid Microsoft Graph v1.0 $search query."
- " The output is a JSON list of the requested resource."
- )
-
- model_config = ConfigDict(
- extra="forbid",
- )
-
- def _run(
- self,
- query: str,
- folder: str = "",
- max_results: int = 10,
- truncate: bool = True,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- truncate_limit: int = 150,
- ) -> List[Dict[str, Any]]:
- # Get mailbox object
- mailbox = self.account.mailbox()
-
- # Pull the folder if the user wants to search in a folder
- if folder != "":
- mailbox = mailbox.get_folder(folder_name=folder)
-
- # Retrieve messages based on query
- query = mailbox.q().search(query)
- messages = mailbox.get_messages(limit=max_results, query=query)
-
- # Generate output dict
- output_messages = []
- for message in messages:
- output_message = {}
- output_message["from"] = message.sender
-
- if truncate:
- output_message["body"] = message.body_preview[:truncate_limit]
- else:
- output_message["body"] = clean_body(message.body)
-
- output_message["subject"] = message.subject
-
- output_message["date"] = message.modified.strftime(UTC_FORMAT)
-
- output_message["to"] = []
- for recipient in message.to._recipients:
- output_message["to"].append(str(recipient))
-
- output_message["cc"] = []
- for recipient in message.cc._recipients:
- output_message["cc"].append(str(recipient))
-
- output_message["bcc"] = []
- for recipient in message.bcc._recipients:
- output_message["bcc"].append(str(recipient))
-
- output_messages.append(output_message)
-
- return output_messages
diff --git a/libs/community/langchain_community/tools/office365/send_event.py b/libs/community/langchain_community/tools/office365/send_event.py
deleted file mode 100644
index 2ab140ca46..0000000000
--- a/libs/community/langchain_community/tools/office365/send_event.py
+++ /dev/null
@@ -1,96 +0,0 @@
-"""Util that sends calendar events in Office 365.
-
-Free, but setup is required. See link below.
-https://learn.microsoft.com/en-us/graph/auth/
-"""
-
-from datetime import datetime as dt
-from typing import List, Optional, Type
-from zoneinfo import ZoneInfo
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.office365.base import O365BaseTool
-from langchain_community.tools.office365.utils import UTC_FORMAT
-
-
-class SendEventSchema(BaseModel):
- """Input for CreateEvent Tool."""
-
- body: str = Field(
- ...,
- description="The message body to include in the event.",
- )
- attendees: List[str] = Field(
- ...,
- description="The list of attendees for the event.",
- )
- subject: str = Field(
- ...,
- description="The subject of the event.",
- )
- start_datetime: str = Field(
- description=" The start datetime for the event in the following format: "
- ' YYYY-MM-DDTHH:MM:SS±hh:mm, where "T" separates the date and time '
- " components, and the time zone offset is specified as ±hh:mm. "
- ' For example: "2023-06-09T10:30:00+03:00" represents June 9th, '
- " 2023, at 10:30 AM in a time zone with a positive offset of 3 "
- " hours from Coordinated Universal Time (UTC).",
- )
- end_datetime: str = Field(
- description=" The end datetime for the event in the following format: "
- ' YYYY-MM-DDTHH:MM:SS±hh:mm, where "T" separates the date and time '
- " components, and the time zone offset is specified as ±hh:mm. "
- ' For example: "2023-06-09T10:30:00+03:00" represents June 9th, '
- " 2023, at 10:30 AM in a time zone with a positive offset of 3 "
- " hours from Coordinated Universal Time (UTC).",
- )
-
-
-class O365SendEvent(O365BaseTool):
- """Tool for sending calendar events in Office 365."""
-
- name: str = "send_event"
- description: str = (
- "Use this tool to create and send an event with the provided event fields."
- )
- args_schema: Type[SendEventSchema] = SendEventSchema
-
- def _run(
- self,
- body: str,
- attendees: List[str],
- subject: str,
- start_datetime: str,
- end_datetime: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- # Get calendar object
- schedule = self.account.schedule()
- calendar = schedule.get_default_calendar()
-
- event = calendar.new_event()
-
- event.body = body
- event.subject = subject
- try:
- event.start = dt.fromisoformat(start_datetime).replace(
- tzinfo=ZoneInfo("UTC")
- )
- except ValueError:
- # fallback for backwards compatibility
- event.start = dt.strptime(start_datetime, UTC_FORMAT)
- try:
- event.end = dt.fromisoformat(end_datetime).replace(tzinfo=ZoneInfo("UTC"))
- except ValueError:
- # fallback for backwards compatibility
- event.end = dt.strptime(end_datetime, UTC_FORMAT)
-
- for attendee in attendees:
- event.attendees.add(attendee)
-
- event.save()
-
- output = "Event sent: " + str(event)
- return output
diff --git a/libs/community/langchain_community/tools/office365/send_message.py b/libs/community/langchain_community/tools/office365/send_message.py
deleted file mode 100644
index 6ebc888371..0000000000
--- a/libs/community/langchain_community/tools/office365/send_message.py
+++ /dev/null
@@ -1,68 +0,0 @@
-from typing import List, Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.office365.base import O365BaseTool
-
-
-class SendMessageSchema(BaseModel):
- """Input for SendMessageTool."""
-
- body: str = Field(
- ...,
- description="The message body to be sent.",
- )
- to: List[str] = Field(
- ...,
- description="The list of recipients.",
- )
- subject: str = Field(
- ...,
- description="The subject of the message.",
- )
- cc: Optional[List[str]] = Field(
- None,
- description="The list of CC recipients.",
- )
- bcc: Optional[List[str]] = Field(
- None,
- description="The list of BCC recipients.",
- )
-
-
-class O365SendMessage(O365BaseTool):
- """Send an email in Office 365."""
-
- name: str = "send_email"
- description: str = (
- "Use this tool to send an email with the provided message fields."
- )
- args_schema: Type[SendMessageSchema] = SendMessageSchema
-
- def _run(
- self,
- body: str,
- to: List[str],
- subject: str,
- cc: Optional[List[str]] = None,
- bcc: Optional[List[str]] = None,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- # Get mailbox object
- mailbox = self.account.mailbox()
- message = mailbox.new_message()
-
- # Assign message values
- message.body = body
- message.subject = subject
- message.to.add(to)
- if cc is not None:
- message.cc.add(cc)
- if bcc is not None:
- message.bcc.add(bcc)
-
- message.send()
-
- output = "Message sent: " + str(message)
- return output
diff --git a/libs/community/langchain_community/tools/office365/utils.py b/libs/community/langchain_community/tools/office365/utils.py
deleted file mode 100644
index 168fe1ccb4..0000000000
--- a/libs/community/langchain_community/tools/office365/utils.py
+++ /dev/null
@@ -1,79 +0,0 @@
-"""O365 tool utils."""
-
-from __future__ import annotations
-
-import logging
-import os
-from typing import TYPE_CHECKING
-
-if TYPE_CHECKING:
- from O365 import Account
-
-logger = logging.getLogger(__name__)
-
-
-def clean_body(body: str) -> str:
- """Clean body of a message or event."""
- try:
- from bs4 import BeautifulSoup
-
- try:
- # Remove HTML
- soup = BeautifulSoup(str(body), "html.parser")
- body = soup.get_text()
-
- # Remove return characters
- body = "".join(body.splitlines())
-
- # Remove extra spaces
- body = " ".join(body.split())
-
- return str(body)
- except Exception:
- return str(body)
- except ImportError:
- return str(body)
-
-
-def authenticate() -> Account:
- """Authenticate using the Microsoft Graph API"""
- try:
- from O365 import Account
- except ImportError as e:
- raise ImportError(
- "Cannot import 0365. Please install the package with `pip install O365`."
- ) from e
-
- if "CLIENT_ID" in os.environ and "CLIENT_SECRET" in os.environ:
- client_id = os.environ["CLIENT_ID"]
- client_secret = os.environ["CLIENT_SECRET"]
- credentials = (client_id, client_secret)
- else:
- logger.error(
- "Error: The CLIENT_ID and CLIENT_SECRET environmental variables have not "
- "been set. Visit the following link on how to acquire these authorization "
- "tokens: https://learn.microsoft.com/en-us/graph/auth/"
- )
- return None
-
- account = Account(credentials)
-
- if account.is_authenticated is False:
- if not account.authenticate(
- scopes=[
- "https://graph.microsoft.com/Mail.ReadWrite",
- "https://graph.microsoft.com/Mail.Send",
- "https://graph.microsoft.com/Calendars.ReadWrite",
- "https://graph.microsoft.com/MailboxSettings.ReadWrite",
- ]
- ):
- print("Error: Could not authenticate") # noqa: T201
- return None
- else:
- return account
- else:
- return account
-
-
-UTC_FORMAT = "%Y-%m-%dT%H:%M:%S%z"
-"""UTC format for datetime objects."""
diff --git a/libs/community/langchain_community/tools/openai_dalle_image_generation/__init__.py b/libs/community/langchain_community/tools/openai_dalle_image_generation/__init__.py
deleted file mode 100644
index dbdf41b11b..0000000000
--- a/libs/community/langchain_community/tools/openai_dalle_image_generation/__init__.py
+++ /dev/null
@@ -1,7 +0,0 @@
-"""Tool to generate an image using DALLE OpenAI V1 SDK."""
-
-from langchain_community.tools.openai_dalle_image_generation.tool import (
- OpenAIDALLEImageGenerationTool,
-)
-
-__all__ = ["OpenAIDALLEImageGenerationTool"]
diff --git a/libs/community/langchain_community/tools/openai_dalle_image_generation/tool.py b/libs/community/langchain_community/tools/openai_dalle_image_generation/tool.py
deleted file mode 100644
index 36374e887f..0000000000
--- a/libs/community/langchain_community/tools/openai_dalle_image_generation/tool.py
+++ /dev/null
@@ -1,29 +0,0 @@
-"""Tool for the OpenAI DALLE V1 Image Generation SDK."""
-
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-
-from langchain_community.utilities.dalle_image_generator import DallEAPIWrapper
-
-
-class OpenAIDALLEImageGenerationTool(BaseTool):
- """Tool that generates an image using OpenAI DALLE."""
-
- name: str = "openai_dalle"
- description: str = (
- "A wrapper around OpenAI DALLE Image Generation. "
- "Useful for when you need to generate an image of"
- "people, places, paintings, animals, or other subjects. "
- "Input should be a text prompt to generate an image."
- )
- api_wrapper: DallEAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the OpenAI DALLE Image Generation tool."""
- return self.api_wrapper.run(query)
diff --git a/libs/community/langchain_community/tools/openapi/__init__.py b/libs/community/langchain_community/tools/openapi/__init__.py
deleted file mode 100644
index e69de29bb2..0000000000
diff --git a/libs/community/langchain_community/tools/openapi/utils/__init__.py b/libs/community/langchain_community/tools/openapi/utils/__init__.py
deleted file mode 100644
index e69de29bb2..0000000000
diff --git a/libs/community/langchain_community/tools/openapi/utils/api_models.py b/libs/community/langchain_community/tools/openapi/utils/api_models.py
deleted file mode 100644
index 8358305464..0000000000
--- a/libs/community/langchain_community/tools/openapi/utils/api_models.py
+++ /dev/null
@@ -1,632 +0,0 @@
-"""Pydantic models for parsing an OpenAPI spec."""
-
-from __future__ import annotations
-
-import logging
-from enum import Enum
-from typing import (
- TYPE_CHECKING,
- Any,
- Dict,
- List,
- Optional,
- Sequence,
- Tuple,
- Type,
- Union,
-)
-
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.openapi.utils.openapi_utils import HTTPVerb, OpenAPISpec
-
-logger = logging.getLogger(__name__)
-PRIMITIVE_TYPES = {
- "integer": int,
- "number": float,
- "string": str,
- "boolean": bool,
- "array": List,
- "object": Dict,
- "null": None,
-}
-
-
-# See https://github.com/OAI/OpenAPI-Specification/blob/main/versions/3.1.0.md#parameterIn
-# for more info.
-class APIPropertyLocation(Enum):
- """The location of the property."""
-
- QUERY = "query"
- PATH = "path"
- HEADER = "header"
- COOKIE = "cookie" # Not yet supported
-
- @classmethod
- def from_str(cls, location: str) -> "APIPropertyLocation":
- """Parse an APIPropertyLocation."""
- try:
- return cls(location)
- except ValueError:
- raise ValueError(
- f"Invalid APIPropertyLocation. Valid values are {cls.__members__}"
- )
-
-
-_SUPPORTED_MEDIA_TYPES = ("application/json",)
-
-SUPPORTED_LOCATIONS = {
- APIPropertyLocation.HEADER,
- APIPropertyLocation.QUERY,
- APIPropertyLocation.PATH,
-}
-INVALID_LOCATION_TEMPL = (
- 'Unsupported APIPropertyLocation "{location}"'
- " for parameter {name}. "
- + f"Valid values are {[loc.value for loc in SUPPORTED_LOCATIONS]}"
-)
-
-SCHEMA_TYPE = Union[str, Type, tuple, None, Enum]
-
-
-class APIPropertyBase(BaseModel):
- """Base model for an API property."""
-
- # The name of the parameter is required and is case-sensitive.
- # If "in" is "path", the "name" field must correspond to a template expression
- # within the path field in the Paths Object.
- # If "in" is "header" and the "name" field is "Accept", "Content-Type",
- # or "Authorization", the parameter definition is ignored.
- # For all other cases, the "name" corresponds to the parameter
- # name used by the "in" property.
- name: str = Field(alias="name")
- """The name of the property."""
-
- required: bool = Field(alias="required")
- """Whether the property is required."""
-
- type: SCHEMA_TYPE = Field(alias="type")
- """The type of the property.
-
- Either a primitive type, a component/parameter type,
- or an array or 'object' (dict) of the above."""
-
- default: Optional[Any] = Field(alias="default", default=None)
- """The default value of the property."""
-
- description: Optional[str] = Field(alias="description", default=None)
- """The description of the property."""
-
-
-if TYPE_CHECKING:
- from openapi_pydantic import (
- MediaType,
- Parameter,
- RequestBody,
- Schema,
- )
-
-
-class APIProperty(APIPropertyBase):
- """A model for a property in the query, path, header, or cookie params."""
-
- location: APIPropertyLocation = Field(alias="location")
- """The path/how it's being passed to the endpoint."""
-
- @staticmethod
- def _cast_schema_list_type(
- schema: Schema,
- ) -> Optional[Union[str, Tuple[str, ...]]]:
- type_ = schema.type
- if not isinstance(type_, list):
- return type_
- else:
- return tuple(type_)
-
- @staticmethod
- def _get_schema_type_for_enum(parameter: Parameter, schema: Schema) -> Enum:
- """Get the schema type when the parameter is an enum."""
- param_name = f"{parameter.name}Enum"
- return Enum(param_name, {str(v): v for v in schema.enum})
-
- @staticmethod
- def _get_schema_type_for_array(
- schema: Schema,
- ) -> Optional[Union[str, Tuple[str, ...]]]:
- from openapi_pydantic import (
- Reference,
- Schema,
- )
-
- items = schema.items
- if isinstance(items, Schema):
- schema_type = APIProperty._cast_schema_list_type(items)
- elif isinstance(items, Reference):
- ref_name = items.ref.split("/")[-1]
- schema_type = ref_name # TODO: Add ref definitions to make his valid
- else:
- raise ValueError(f"Unsupported array items: {items}")
-
- if isinstance(schema_type, str):
- # TODO: recurse
- schema_type = (schema_type,)
-
- return schema_type
-
- @staticmethod
- def _get_schema_type(parameter: Parameter, schema: Optional[Schema]) -> SCHEMA_TYPE:
- if schema is None:
- return None
- schema_type: SCHEMA_TYPE = APIProperty._cast_schema_list_type(schema)
- if schema_type == "array":
- schema_type = APIProperty._get_schema_type_for_array(schema)
- elif schema_type == "object":
- # TODO: Resolve array and object types to components.
- raise NotImplementedError("Objects not yet supported")
- elif schema_type in PRIMITIVE_TYPES:
- if schema.enum:
- schema_type = APIProperty._get_schema_type_for_enum(parameter, schema)
- else:
- # Directly use the primitive type
- pass
- else:
- raise NotImplementedError(f"Unsupported type: {schema_type}")
-
- return schema_type
-
- @staticmethod
- def _validate_location(location: APIPropertyLocation, name: str) -> None:
- if location not in SUPPORTED_LOCATIONS:
- raise NotImplementedError(
- INVALID_LOCATION_TEMPL.format(location=location, name=name)
- )
-
- @staticmethod
- def _validate_content(content: Optional[Dict[str, MediaType]]) -> None:
- if content:
- raise ValueError(
- "API Properties with media content not supported. "
- "Media content only supported within APIRequestBodyProperty's"
- )
-
- @staticmethod
- def _get_schema(parameter: Parameter, spec: OpenAPISpec) -> Optional[Schema]:
- from openapi_pydantic import (
- Reference,
- Schema,
- )
-
- schema = parameter.param_schema
- if isinstance(schema, Reference):
- schema = spec.get_referenced_schema(schema)
- elif schema is None:
- return None
- elif not isinstance(schema, Schema):
- raise ValueError(f"Error dereferencing schema: {schema}")
-
- return schema
-
- @staticmethod
- def is_supported_location(location: str) -> bool:
- """Return whether the provided location is supported."""
- try:
- return APIPropertyLocation.from_str(location) in SUPPORTED_LOCATIONS
- except ValueError:
- return False
-
- @classmethod
- def from_parameter(cls, parameter: Parameter, spec: OpenAPISpec) -> "APIProperty":
- """Instantiate from an OpenAPI Parameter."""
- location = APIPropertyLocation.from_str(parameter.param_in)
- cls._validate_location(
- location,
- parameter.name,
- )
- cls._validate_content(parameter.content)
- schema = cls._get_schema(parameter, spec)
- schema_type = cls._get_schema_type(parameter, schema)
- default_val = schema.default if schema is not None else None
- return cls(
- name=parameter.name,
- location=location,
- default=default_val,
- description=parameter.description,
- required=parameter.required,
- type=schema_type,
- )
-
-
-class APIRequestBodyProperty(APIPropertyBase):
- """A model for a request body property."""
-
- properties: List["APIRequestBodyProperty"] = Field(alias="properties")
- """The sub-properties of the property."""
-
- # This is useful for handling nested property cycles.
- # We can define separate types in that case.
- references_used: List[str] = Field(alias="references_used")
- """The references used by the property."""
-
- @classmethod
- def _process_object_schema(
- cls, schema: Schema, spec: OpenAPISpec, references_used: List[str]
- ) -> Tuple[Union[str, List[str], None], List["APIRequestBodyProperty"]]:
- from openapi_pydantic import (
- Reference,
- )
-
- properties = []
- required_props = schema.required or []
- if schema.properties is None:
- raise ValueError(
- f"No properties found when processing object schema: {schema}"
- )
- for prop_name, prop_schema in schema.properties.items():
- if isinstance(prop_schema, Reference):
- ref_name = prop_schema.ref.split("/")[-1]
- if ref_name not in references_used:
- references_used.append(ref_name)
- prop_schema = spec.get_referenced_schema(prop_schema)
- else:
- continue
-
- properties.append(
- cls.from_schema(
- schema=prop_schema,
- name=prop_name,
- required=prop_name in required_props,
- spec=spec,
- references_used=references_used,
- )
- )
- return schema.type, properties
-
- @classmethod
- def _process_array_schema(
- cls,
- schema: Schema,
- name: str,
- spec: OpenAPISpec,
- references_used: List[str],
- ) -> str:
- from openapi_pydantic import Reference, Schema
-
- items = schema.items
- if items is not None:
- if isinstance(items, Reference):
- ref_name = items.ref.split("/")[-1]
- if ref_name not in references_used:
- references_used.append(ref_name)
- items = spec.get_referenced_schema(items)
- else:
- pass
- return f"Array<{ref_name}>"
- else:
- pass
-
- if isinstance(items, Schema):
- array_type = cls.from_schema(
- schema=items,
- name=f"{name}Item",
- required=True, # TODO: Add required
- spec=spec,
- references_used=references_used,
- )
- return f"Array<{array_type.type}>"
-
- return "array"
-
- @classmethod
- def from_schema(
- cls,
- schema: Schema,
- name: str,
- required: bool,
- spec: OpenAPISpec,
- references_used: Optional[List[str]] = None,
- ) -> "APIRequestBodyProperty":
- """Recursively populate from an OpenAPI Schema."""
- if references_used is None:
- references_used = []
-
- schema_type = schema.type
- properties: List[APIRequestBodyProperty] = []
- if schema_type == "object" and schema.properties:
- schema_type, properties = cls._process_object_schema(
- schema, spec, references_used
- )
- elif schema_type == "array":
- schema_type = cls._process_array_schema(schema, name, spec, references_used)
- elif schema_type in PRIMITIVE_TYPES:
- # Use the primitive type directly
- pass
- elif schema_type is None:
- # No typing specified/parsed. WIll map to 'any'
- pass
- else:
- raise ValueError(f"Unsupported type: {schema_type}")
-
- return cls(
- name=name,
- required=required,
- type=schema_type,
- default=schema.default,
- description=schema.description,
- properties=properties,
- references_used=references_used,
- )
-
-
-# class APIRequestBodyProperty(APIPropertyBase):
-class APIRequestBody(BaseModel):
- """A model for a request body."""
-
- description: Optional[str] = Field(alias="description")
- """The description of the request body."""
-
- properties: List[APIRequestBodyProperty] = Field(alias="properties")
-
- # E.g., application/json - we only support JSON at the moment.
- media_type: str = Field(alias="media_type")
- """The media type of the request body."""
-
- @classmethod
- def _process_supported_media_type(
- cls,
- media_type_obj: MediaType,
- spec: OpenAPISpec,
- ) -> List[APIRequestBodyProperty]:
- """Process the media type of the request body."""
- from openapi_pydantic import Reference
-
- references_used = []
- schema = media_type_obj.media_type_schema
- if isinstance(schema, Reference):
- references_used.append(schema.ref.split("/")[-1])
- schema = spec.get_referenced_schema(schema)
- if schema is None:
- raise ValueError(
- f"Could not resolve schema for media type: {media_type_obj}"
- )
- api_request_body_properties = []
- required_properties = schema.required or []
- if schema.type == "object" and schema.properties:
- for prop_name, prop_schema in schema.properties.items():
- if isinstance(prop_schema, Reference):
- prop_schema = spec.get_referenced_schema(prop_schema)
-
- api_request_body_properties.append(
- APIRequestBodyProperty.from_schema(
- schema=prop_schema,
- name=prop_name,
- required=prop_name in required_properties,
- spec=spec,
- )
- )
- else:
- api_request_body_properties.append(
- APIRequestBodyProperty(
- name="body",
- required=True,
- type=schema.type,
- default=schema.default,
- description=schema.description,
- properties=[],
- references_used=references_used,
- )
- )
-
- return api_request_body_properties
-
- @classmethod
- def from_request_body(
- cls, request_body: RequestBody, spec: OpenAPISpec
- ) -> "APIRequestBody":
- """Instantiate from an OpenAPI RequestBody."""
- properties = []
- for media_type, media_type_obj in request_body.content.items():
- if media_type not in _SUPPORTED_MEDIA_TYPES:
- continue
- api_request_body_properties = cls._process_supported_media_type(
- media_type_obj,
- spec,
- )
- properties.extend(api_request_body_properties)
-
- return cls(
- description=request_body.description,
- properties=properties,
- media_type=media_type,
- )
-
-
-# class APIRequestBodyProperty(APIPropertyBase):
-# class APIRequestBody(BaseModel):
-class APIOperation(BaseModel):
- """A model for a single API operation."""
-
- operation_id: str = Field(alias="operation_id")
- """The unique identifier of the operation."""
-
- description: Optional[str] = Field(alias="description")
- """The description of the operation."""
-
- base_url: str = Field(alias="base_url")
- """The base URL of the operation."""
-
- path: str = Field(alias="path")
- """The path of the operation."""
-
- method: HTTPVerb = Field(alias="method")
- """The HTTP method of the operation."""
-
- properties: Sequence[APIProperty] = Field(alias="properties")
-
- # TODO: Add parse in used components to be able to specify what type of
- # referenced object it is.
- # """The properties of the operation."""
- # components: Dict[str, BaseModel] = Field(alias="components")
-
- request_body: Optional[APIRequestBody] = Field(alias="request_body")
- """The request body of the operation."""
-
- @staticmethod
- def _get_properties_from_parameters(
- parameters: List[Parameter], spec: OpenAPISpec
- ) -> List[APIProperty]:
- """Get the properties of the operation."""
- properties = []
- for param in parameters:
- if APIProperty.is_supported_location(param.param_in):
- properties.append(APIProperty.from_parameter(param, spec))
- elif param.required:
- raise ValueError(
- INVALID_LOCATION_TEMPL.format(
- location=param.param_in, name=param.name
- )
- )
- else:
- logger.warning(
- INVALID_LOCATION_TEMPL.format(
- location=param.param_in, name=param.name
- )
- + " Ignoring optional parameter"
- )
- pass
- return properties
-
- @classmethod
- def from_openapi_url(
- cls,
- spec_url: str,
- path: str,
- method: str,
- ) -> "APIOperation":
- """Create an APIOperation from an OpenAPI URL."""
- spec = OpenAPISpec.from_url(spec_url)
- return cls.from_openapi_spec(spec, path, method)
-
- @classmethod
- def from_openapi_spec(
- cls,
- spec: OpenAPISpec,
- path: str,
- method: str,
- ) -> "APIOperation":
- """Create an APIOperation from an OpenAPI spec."""
- operation = spec.get_operation(path, method)
- parameters = spec.get_parameters_for_operation(operation)
- properties = cls._get_properties_from_parameters(parameters, spec)
- operation_id = OpenAPISpec.get_cleaned_operation_id(operation, path, method)
- request_body = spec.get_request_body_for_operation(operation)
- api_request_body = (
- APIRequestBody.from_request_body(request_body, spec)
- if request_body is not None
- else None
- )
- description = operation.description or operation.summary
- if not description and spec.paths is not None:
- description = spec.paths[path].description or spec.paths[path].summary
- return cls(
- operation_id=operation_id,
- description=description or "",
- base_url=spec.base_url,
- path=path,
- method=method, # type: ignore[arg-type]
- properties=properties,
- request_body=api_request_body,
- )
-
- @staticmethod
- def ts_type_from_python(type_: SCHEMA_TYPE) -> str:
- if type_ is None:
- # TODO: Handle Nones better. These often result when
- # parsing specs that are < v3
- return "any"
- elif isinstance(type_, str):
- return {
- "str": "string",
- "integer": "number",
- "float": "number",
- "date-time": "string",
- }.get(type_, type_)
- elif isinstance(type_, tuple):
- return f"Array<{APIOperation.ts_type_from_python(type_[0])}>"
- elif isinstance(type_, type) and issubclass(type_, Enum):
- return " | ".join([f"'{e.value}'" for e in type_])
- else:
- return str(type_)
-
- def _format_nested_properties(
- self, properties: List[APIRequestBodyProperty], indent: int = 2
- ) -> str:
- """Format nested properties."""
- formatted_props = []
-
- for prop in properties:
- prop_name = prop.name
- prop_type = self.ts_type_from_python(prop.type)
- prop_required = "" if prop.required else "?"
- prop_desc = f"/* {prop.description} */" if prop.description else ""
-
- if prop.properties:
- nested_props = self._format_nested_properties(
- prop.properties, indent + 2
- )
- prop_type = f"{{\n{nested_props}\n{' ' * indent}}}"
-
- formatted_props.append(
- f"{prop_desc}\n{' ' * indent}{prop_name}{prop_required}: {prop_type},"
- )
-
- return "\n".join(formatted_props)
-
- def to_typescript(self) -> str:
- """Get typescript string representation of the operation."""
- operation_name = self.operation_id
- params = []
-
- if self.request_body:
- formatted_request_body_props = self._format_nested_properties(
- self.request_body.properties
- )
- params.append(formatted_request_body_props)
-
- for prop in self.properties:
- prop_name = prop.name
- prop_type = self.ts_type_from_python(prop.type)
- prop_required = "" if prop.required else "?"
- prop_desc = f"/* {prop.description} */" if prop.description else ""
- params.append(f"{prop_desc}\n\t\t{prop_name}{prop_required}: {prop_type},")
-
- formatted_params = "\n".join(params).strip()
- description_str = f"/* {self.description} */" if self.description else ""
- typescript_definition = f"""
-{description_str}
-type {operation_name} = (_: {{
-{formatted_params}
-}}) => any;
-"""
- return typescript_definition.strip()
-
- @property
- def query_params(self) -> List[str]:
- return [
- property.name
- for property in self.properties
- if property.location == APIPropertyLocation.QUERY
- ]
-
- @property
- def path_params(self) -> List[str]:
- return [
- property.name
- for property in self.properties
- if property.location == APIPropertyLocation.PATH
- ]
-
- @property
- def body_params(self) -> List[str]:
- if self.request_body is None:
- return []
- return [prop.name for prop in self.request_body.properties]
diff --git a/libs/community/langchain_community/tools/openapi/utils/openapi_utils.py b/libs/community/langchain_community/tools/openapi/utils/openapi_utils.py
deleted file mode 100644
index 7ed0ade18d..0000000000
--- a/libs/community/langchain_community/tools/openapi/utils/openapi_utils.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""Utility functions for parsing an OpenAPI spec. Kept for backwards compat."""
-
-from langchain_community.utilities.openapi import HTTPVerb, OpenAPISpec
-
-__all__ = ["HTTPVerb", "OpenAPISpec"]
diff --git a/libs/community/langchain_community/tools/openweathermap/__init__.py b/libs/community/langchain_community/tools/openweathermap/__init__.py
deleted file mode 100644
index eb9abd3ccd..0000000000
--- a/libs/community/langchain_community/tools/openweathermap/__init__.py
+++ /dev/null
@@ -1,7 +0,0 @@
-"""OpenWeatherMap API toolkit."""
-
-from langchain_community.tools.openweathermap.tool import OpenWeatherMapQueryRun
-
-__all__ = [
- "OpenWeatherMapQueryRun",
-]
diff --git a/libs/community/langchain_community/tools/openweathermap/tool.py b/libs/community/langchain_community/tools/openweathermap/tool.py
deleted file mode 100644
index f88095d3ef..0000000000
--- a/libs/community/langchain_community/tools/openweathermap/tool.py
+++ /dev/null
@@ -1,30 +0,0 @@
-"""Tool for the OpenWeatherMap API."""
-
-from typing import Optional
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import Field
-
-from langchain_community.utilities.openweathermap import OpenWeatherMapAPIWrapper
-
-
-class OpenWeatherMapQueryRun(BaseTool):
- """Tool that queries the OpenWeatherMap API."""
-
- api_wrapper: OpenWeatherMapAPIWrapper = Field(
- default_factory=OpenWeatherMapAPIWrapper
- )
-
- name: str = "open_weather_map"
- description: str = (
- "A wrapper around OpenWeatherMap API. "
- "Useful for fetching current weather information for a specified location. "
- "Input should be a location string (e.g. London,GB)."
- )
-
- def _run(
- self, location: str, run_manager: Optional[CallbackManagerForToolRun] = None
- ) -> str:
- """Use the OpenWeatherMap tool."""
- return self.api_wrapper.run(location)
diff --git a/libs/community/langchain_community/tools/passio_nutrition_ai/__init__.py b/libs/community/langchain_community/tools/passio_nutrition_ai/__init__.py
deleted file mode 100644
index f75469d3f1..0000000000
--- a/libs/community/langchain_community/tools/passio_nutrition_ai/__init__.py
+++ /dev/null
@@ -1,5 +0,0 @@
-"""Passio Nutrition AI API toolkit."""
-
-from langchain_community.tools.passio_nutrition_ai.tool import NutritionAI
-
-__all__ = ["NutritionAI"]
diff --git a/libs/community/langchain_community/tools/passio_nutrition_ai/tool.py b/libs/community/langchain_community/tools/passio_nutrition_ai/tool.py
deleted file mode 100644
index 939e1a41bc..0000000000
--- a/libs/community/langchain_community/tools/passio_nutrition_ai/tool.py
+++ /dev/null
@@ -1,38 +0,0 @@
-"""Tool for the Passio Nutrition AI API."""
-
-from typing import Dict, Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.utilities.passio_nutrition_ai import NutritionAIAPI
-
-
-class NutritionAIInputs(BaseModel):
- """Inputs to the Passio Nutrition AI tool."""
-
- query: str = Field(
- description="A query to look up using Passio Nutrition AI, usually a few words."
- )
-
-
-class NutritionAI(BaseTool):
- """Tool that queries the Passio Nutrition AI API."""
-
- name: str = "nutritionai_advanced_search"
- description: str = (
- "A wrapper around the Passio Nutrition AI. "
- "Useful to retrieve nutrition facts. "
- "Input should be a search query string."
- )
- api_wrapper: NutritionAIAPI
- args_schema: Type[BaseModel] = NutritionAIInputs
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> Optional[Dict]:
- """Use the tool."""
- return self.api_wrapper.run(query)
diff --git a/libs/community/langchain_community/tools/playwright/__init__.py b/libs/community/langchain_community/tools/playwright/__init__.py
deleted file mode 100644
index f69ff8025d..0000000000
--- a/libs/community/langchain_community/tools/playwright/__init__.py
+++ /dev/null
@@ -1,21 +0,0 @@
-"""Browser tools and toolkit."""
-
-from langchain_community.tools.playwright.click import ClickTool
-from langchain_community.tools.playwright.current_page import CurrentWebPageTool
-from langchain_community.tools.playwright.extract_hyperlinks import (
- ExtractHyperlinksTool,
-)
-from langchain_community.tools.playwright.extract_text import ExtractTextTool
-from langchain_community.tools.playwright.get_elements import GetElementsTool
-from langchain_community.tools.playwright.navigate import NavigateTool
-from langchain_community.tools.playwright.navigate_back import NavigateBackTool
-
-__all__ = [
- "NavigateTool",
- "NavigateBackTool",
- "ExtractTextTool",
- "ExtractHyperlinksTool",
- "GetElementsTool",
- "ClickTool",
- "CurrentWebPageTool",
-]
diff --git a/libs/community/langchain_community/tools/playwright/base.py b/libs/community/langchain_community/tools/playwright/base.py
deleted file mode 100644
index e85cc84793..0000000000
--- a/libs/community/langchain_community/tools/playwright/base.py
+++ /dev/null
@@ -1,58 +0,0 @@
-from __future__ import annotations
-
-from typing import TYPE_CHECKING, Any, Optional, Tuple, Type
-
-from langchain_core.tools import BaseTool
-from langchain_core.utils import guard_import
-from pydantic import model_validator
-
-if TYPE_CHECKING:
- from playwright.async_api import Browser as AsyncBrowser
- from playwright.sync_api import Browser as SyncBrowser
-else:
- try:
- # We do this so pydantic can resolve the types when instantiating
- from playwright.async_api import Browser as AsyncBrowser
- from playwright.sync_api import Browser as SyncBrowser
- except ImportError:
- pass
-
-
-def lazy_import_playwright_browsers() -> Tuple[Type[AsyncBrowser], Type[SyncBrowser]]:
- """
- Lazy import playwright browsers.
-
- Returns:
- Tuple[Type[AsyncBrowser], Type[SyncBrowser]]:
- AsyncBrowser and SyncBrowser classes.
- """
- return (
- guard_import(module_name="playwright.async_api").Browser,
- guard_import(module_name="playwright.sync_api").Browser,
- )
-
-
-class BaseBrowserTool(BaseTool):
- """Base class for browser tools."""
-
- sync_browser: Optional["SyncBrowser"] = None
- async_browser: Optional["AsyncBrowser"] = None
-
- @model_validator(mode="before")
- @classmethod
- def validate_browser_provided(cls, values: dict) -> Any:
- """Check that the arguments are valid."""
- lazy_import_playwright_browsers()
- if values.get("async_browser") is None and values.get("sync_browser") is None:
- raise ValueError("Either async_browser or sync_browser must be specified.")
- return values
-
- @classmethod
- def from_browser(
- cls,
- sync_browser: Optional[SyncBrowser] = None,
- async_browser: Optional[AsyncBrowser] = None,
- ) -> BaseBrowserTool:
- """Instantiate the tool."""
- lazy_import_playwright_browsers()
- return cls(sync_browser=sync_browser, async_browser=async_browser) # type: ignore[call-arg]
diff --git a/libs/community/langchain_community/tools/playwright/click.py b/libs/community/langchain_community/tools/playwright/click.py
deleted file mode 100644
index 22c6a23bf9..0000000000
--- a/libs/community/langchain_community/tools/playwright/click.py
+++ /dev/null
@@ -1,87 +0,0 @@
-from __future__ import annotations
-
-from typing import Optional, Type
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.playwright.base import BaseBrowserTool
-from langchain_community.tools.playwright.utils import (
- aget_current_page,
- get_current_page,
-)
-
-
-class ClickToolInput(BaseModel):
- """Input for ClickTool."""
-
- selector: str = Field(..., description="CSS selector for the element to click")
-
-
-class ClickTool(BaseBrowserTool):
- """Tool for clicking on an element with the given CSS selector."""
-
- name: str = "click_element"
- description: str = "Click on an element with the given CSS selector"
- args_schema: Type[BaseModel] = ClickToolInput
-
- visible_only: bool = True
- """Whether to consider only visible elements."""
- playwright_strict: bool = False
- """Whether to employ Playwright's strict mode when clicking on elements."""
- playwright_timeout: float = 1_000
- """Timeout (in ms) for Playwright to wait for element to be ready."""
-
- def _selector_effective(self, selector: str) -> str:
- if not self.visible_only:
- return selector
- return f"{selector} >> visible=1"
-
- def _run(
- self,
- selector: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- if self.sync_browser is None:
- raise ValueError(f"Synchronous browser not provided to {self.name}")
- page = get_current_page(self.sync_browser)
- # Navigate to the desired webpage before using this tool
- selector_effective = self._selector_effective(selector=selector)
- from playwright.sync_api import TimeoutError as PlaywrightTimeoutError
-
- try:
- page.click(
- selector_effective,
- strict=self.playwright_strict,
- timeout=self.playwright_timeout,
- )
- except PlaywrightTimeoutError:
- return f"Unable to click on element '{selector}'"
- return f"Clicked element '{selector}'"
-
- async def _arun(
- self,
- selector: str,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- if self.async_browser is None:
- raise ValueError(f"Asynchronous browser not provided to {self.name}")
- page = await aget_current_page(self.async_browser)
- # Navigate to the desired webpage before using this tool
- selector_effective = self._selector_effective(selector=selector)
- from playwright.async_api import TimeoutError as PlaywrightTimeoutError
-
- try:
- await page.click(
- selector_effective,
- strict=self.playwright_strict,
- timeout=self.playwright_timeout,
- )
- except PlaywrightTimeoutError:
- return f"Unable to click on element '{selector}'"
- return f"Clicked element '{selector}'"
diff --git a/libs/community/langchain_community/tools/playwright/current_page.py b/libs/community/langchain_community/tools/playwright/current_page.py
deleted file mode 100644
index 207cac4b70..0000000000
--- a/libs/community/langchain_community/tools/playwright/current_page.py
+++ /dev/null
@@ -1,47 +0,0 @@
-from __future__ import annotations
-
-from typing import Optional, Type
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from pydantic import BaseModel
-
-from langchain_community.tools.playwright.base import BaseBrowserTool
-from langchain_community.tools.playwright.utils import (
- aget_current_page,
- get_current_page,
-)
-
-
-class CurrentWebPageToolInput(BaseModel):
- """Explicit no-args input for CurrentWebPageTool."""
-
-
-class CurrentWebPageTool(BaseBrowserTool):
- """Tool for getting the URL of the current webpage."""
-
- name: str = "current_webpage"
- description: str = "Returns the URL of the current page"
- args_schema: Type[BaseModel] = CurrentWebPageToolInput
-
- def _run(
- self,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- if self.sync_browser is None:
- raise ValueError(f"Synchronous browser not provided to {self.name}")
- page = get_current_page(self.sync_browser)
- return str(page.url)
-
- async def _arun(
- self,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- if self.async_browser is None:
- raise ValueError(f"Asynchronous browser not provided to {self.name}")
- page = await aget_current_page(self.async_browser)
- return str(page.url)
diff --git a/libs/community/langchain_community/tools/playwright/extract_hyperlinks.py b/libs/community/langchain_community/tools/playwright/extract_hyperlinks.py
deleted file mode 100644
index 00a5e29027..0000000000
--- a/libs/community/langchain_community/tools/playwright/extract_hyperlinks.py
+++ /dev/null
@@ -1,93 +0,0 @@
-from __future__ import annotations
-
-import json
-from typing import TYPE_CHECKING, Any, Optional, Type
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from pydantic import BaseModel, Field, model_validator
-
-from langchain_community.tools.playwright.base import BaseBrowserTool
-from langchain_community.tools.playwright.utils import (
- aget_current_page,
- get_current_page,
-)
-
-if TYPE_CHECKING:
- pass
-
-
-class ExtractHyperlinksToolInput(BaseModel):
- """Input for ExtractHyperlinksTool."""
-
- absolute_urls: bool = Field(
- default=False,
- description="Return absolute URLs instead of relative URLs",
- )
-
-
-class ExtractHyperlinksTool(BaseBrowserTool):
- """Extract all hyperlinks on the page."""
-
- name: str = "extract_hyperlinks"
- description: str = "Extract all hyperlinks on the current webpage"
- args_schema: Type[BaseModel] = ExtractHyperlinksToolInput
-
- @model_validator(mode="before")
- @classmethod
- def check_bs_import(cls, values: dict) -> Any:
- """Check that the arguments are valid."""
- try:
- from bs4 import BeautifulSoup # noqa: F401
- except ImportError:
- raise ImportError(
- "The 'beautifulsoup4' package is required to use this tool."
- " Please install it with 'pip install beautifulsoup4'."
- )
- return values
-
- @staticmethod
- def scrape_page(page: Any, html_content: str, absolute_urls: bool) -> str:
- from urllib.parse import urljoin
-
- from bs4 import BeautifulSoup
-
- # Parse the HTML content with BeautifulSoup
- soup = BeautifulSoup(html_content, "lxml")
-
- # Find all the anchor elements and extract their href attributes
- anchors = soup.find_all("a")
- if absolute_urls:
- base_url = page.url
- links = [urljoin(base_url, anchor.get("href", "")) for anchor in anchors]
- else:
- links = [anchor.get("href", "") for anchor in anchors]
- # Return the list of links as a JSON string. Duplicated link
- # only appears once in the list
- return json.dumps(list(set(links)))
-
- def _run(
- self,
- absolute_urls: bool = False,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- if self.sync_browser is None:
- raise ValueError(f"Synchronous browser not provided to {self.name}")
- page = get_current_page(self.sync_browser)
- html_content = page.content()
- return self.scrape_page(page, html_content, absolute_urls)
-
- async def _arun(
- self,
- absolute_urls: bool = False,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool asynchronously."""
- if self.async_browser is None:
- raise ValueError(f"Asynchronous browser not provided to {self.name}")
- page = await aget_current_page(self.async_browser)
- html_content = await page.content()
- return self.scrape_page(page, html_content, absolute_urls)
diff --git a/libs/community/langchain_community/tools/playwright/extract_text.py b/libs/community/langchain_community/tools/playwright/extract_text.py
deleted file mode 100644
index 7c9ce7f8e1..0000000000
--- a/libs/community/langchain_community/tools/playwright/extract_text.py
+++ /dev/null
@@ -1,73 +0,0 @@
-from __future__ import annotations
-
-from typing import Any, Optional, Type
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from pydantic import BaseModel, model_validator
-
-from langchain_community.tools.playwright.base import BaseBrowserTool
-from langchain_community.tools.playwright.utils import (
- aget_current_page,
- get_current_page,
-)
-
-
-class ExtractTextToolInput(BaseModel):
- """Explicit no-args input for ExtractTextTool."""
-
-
-class ExtractTextTool(BaseBrowserTool):
- """Tool for extracting all the text on the current webpage."""
-
- name: str = "extract_text"
- description: str = "Extract all the text on the current webpage"
- args_schema: Type[BaseModel] = ExtractTextToolInput
-
- @model_validator(mode="before")
- @classmethod
- def check_acheck_bs_importrgs(cls, values: dict) -> Any:
- """Check that the arguments are valid."""
- try:
- from bs4 import BeautifulSoup # noqa: F401
- except ImportError:
- raise ImportError(
- "The 'beautifulsoup4' package is required to use this tool."
- " Please install it with 'pip install beautifulsoup4'."
- )
- return values
-
- def _run(self, run_manager: Optional[CallbackManagerForToolRun] = None) -> str:
- """Use the tool."""
- # Use Beautiful Soup since it's faster than looping through the elements
- from bs4 import BeautifulSoup
-
- if self.sync_browser is None:
- raise ValueError(f"Synchronous browser not provided to {self.name}")
-
- page = get_current_page(self.sync_browser)
- html_content = page.content()
-
- # Parse the HTML content with BeautifulSoup
- soup = BeautifulSoup(html_content, "lxml")
-
- return " ".join(text for text in soup.stripped_strings)
-
- async def _arun(
- self, run_manager: Optional[AsyncCallbackManagerForToolRun] = None
- ) -> str:
- """Use the tool."""
- if self.async_browser is None:
- raise ValueError(f"Asynchronous browser not provided to {self.name}")
- # Use Beautiful Soup since it's faster than looping through the elements
- from bs4 import BeautifulSoup
-
- page = await aget_current_page(self.async_browser)
- html_content = await page.content()
-
- # Parse the HTML content with BeautifulSoup
- soup = BeautifulSoup(html_content, "lxml")
-
- return " ".join(text for text in soup.stripped_strings)
diff --git a/libs/community/langchain_community/tools/playwright/get_elements.py b/libs/community/langchain_community/tools/playwright/get_elements.py
deleted file mode 100644
index 11e43c0169..0000000000
--- a/libs/community/langchain_community/tools/playwright/get_elements.py
+++ /dev/null
@@ -1,111 +0,0 @@
-from __future__ import annotations
-
-import json
-from typing import TYPE_CHECKING, List, Optional, Sequence, Type
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from pydantic import BaseModel, Field
-
-from langchain_community.tools.playwright.base import BaseBrowserTool
-from langchain_community.tools.playwright.utils import (
- aget_current_page,
- get_current_page,
-)
-
-if TYPE_CHECKING:
- from playwright.async_api import Page as AsyncPage
- from playwright.sync_api import Page as SyncPage
-
-
-class GetElementsToolInput(BaseModel):
- """Input for GetElementsTool."""
-
- selector: str = Field(
- ...,
- description="CSS selector, such as '*', 'div', 'p', 'a', #id, .classname",
- )
- attributes: List[str] = Field(
- default_factory=lambda: ["innerText"],
- description="Set of attributes to retrieve for each element",
- )
-
-
-async def _aget_elements(
- page: AsyncPage, selector: str, attributes: Sequence[str]
-) -> List[dict]:
- """Get elements matching the given CSS selector."""
- elements = await page.query_selector_all(selector)
- results = []
- for element in elements:
- result = {}
- for attribute in attributes:
- if attribute == "innerText":
- val: Optional[str] = await element.inner_text()
- else:
- val = await element.get_attribute(attribute)
- if val is not None and val.strip() != "":
- result[attribute] = val
- if result:
- results.append(result)
- return results
-
-
-def _get_elements(
- page: SyncPage, selector: str, attributes: Sequence[str]
-) -> List[dict]:
- """Get elements matching the given CSS selector."""
- elements = page.query_selector_all(selector)
- results = []
- for element in elements:
- result = {}
- for attribute in attributes:
- if attribute == "innerText":
- val: Optional[str] = element.inner_text()
- else:
- val = element.get_attribute(attribute)
- if val is not None and val.strip() != "":
- result[attribute] = val
- if result:
- results.append(result)
- return results
-
-
-class GetElementsTool(BaseBrowserTool):
- """Tool for getting elements in the current web page matching a CSS selector."""
-
- name: str = "get_elements"
- description: str = (
- "Retrieve elements in the current web page matching the given CSS selector"
- )
- args_schema: Type[BaseModel] = GetElementsToolInput
-
- def _run(
- self,
- selector: str,
- attributes: Sequence[str] = ["innerText"],
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- if self.sync_browser is None:
- raise ValueError(f"Synchronous browser not provided to {self.name}")
- page = get_current_page(self.sync_browser)
- # Navigate to the desired webpage before using this tool
- results = _get_elements(page, selector, attributes)
- return json.dumps(results, ensure_ascii=False)
-
- async def _arun(
- self,
- selector: str,
- attributes: Sequence[str] = ["innerText"],
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- if self.async_browser is None:
- raise ValueError(f"Asynchronous browser not provided to {self.name}")
- page = await aget_current_page(self.async_browser)
- # Navigate to the desired webpage before using this tool
- results = await _aget_elements(page, selector, attributes)
- return json.dumps(results, ensure_ascii=False)
diff --git a/libs/community/langchain_community/tools/playwright/navigate.py b/libs/community/langchain_community/tools/playwright/navigate.py
deleted file mode 100644
index 2bfe2be4fd..0000000000
--- a/libs/community/langchain_community/tools/playwright/navigate.py
+++ /dev/null
@@ -1,83 +0,0 @@
-from __future__ import annotations
-
-from typing import Optional, Type
-from urllib.parse import urlparse
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from pydantic import BaseModel, Field, model_validator
-
-from langchain_community.tools.playwright.base import BaseBrowserTool
-from langchain_community.tools.playwright.utils import (
- aget_current_page,
- get_current_page,
-)
-
-
-class NavigateToolInput(BaseModel):
- """Input for NavigateToolInput."""
-
- url: str = Field(..., description="url to navigate to")
-
- @model_validator(mode="before")
- @classmethod
- def validate_url_scheme(cls, values: dict) -> dict:
- """Check that the URL scheme is valid."""
- url = values.get("url")
- parsed_url = urlparse(url)
- if parsed_url.scheme not in ("http", "https"):
- raise ValueError("URL scheme must be 'http' or 'https'")
- return values
-
-
-class NavigateTool(BaseBrowserTool):
- """Tool for navigating a browser to a URL.
-
- **Security Note**: This tool provides code to control web-browser navigation.
-
- This tool can navigate to any URL, including internal network URLs, and
- URLs exposed on the server itself.
-
- However, if exposing this tool to end-users, consider limiting network
- access to the server that hosts the agent.
-
- By default, the URL scheme has been limited to 'http' and 'https' to
- prevent navigation to local file system URLs (or other schemes).
-
- If access to the local file system is required, consider creating a custom
- tool or providing a custom args_schema that allows the desired URL schemes.
-
- See https://python.langchain.com/docs/security for more information.
- """
-
- name: str = "navigate_browser"
- description: str = "Navigate a browser to the specified URL"
- args_schema: Type[BaseModel] = NavigateToolInput
-
- def _run(
- self,
- url: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- if self.sync_browser is None:
- raise ValueError(f"Synchronous browser not provided to {self.name}")
- page = get_current_page(self.sync_browser)
- response = page.goto(url)
- status = response.status if response else "unknown"
- return f"Navigating to {url} returned status code {status}"
-
- async def _arun(
- self,
- url: str,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- if self.async_browser is None:
- raise ValueError(f"Asynchronous browser not provided to {self.name}")
- page = await aget_current_page(self.async_browser)
- response = await page.goto(url)
- status = response.status if response else "unknown"
- return f"Navigating to {url} returned status code {status}"
diff --git a/libs/community/langchain_community/tools/playwright/navigate_back.py b/libs/community/langchain_community/tools/playwright/navigate_back.py
deleted file mode 100644
index 45fa250cb4..0000000000
--- a/libs/community/langchain_community/tools/playwright/navigate_back.py
+++ /dev/null
@@ -1,60 +0,0 @@
-from __future__ import annotations
-
-from typing import Optional, Type
-
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from pydantic import BaseModel
-
-from langchain_community.tools.playwright.base import BaseBrowserTool
-from langchain_community.tools.playwright.utils import (
- aget_current_page,
- get_current_page,
-)
-
-
-class NavigateBackToolInput(BaseModel):
- """Explicit no-args input for NavigateBackTool."""
-
-
-class NavigateBackTool(BaseBrowserTool):
- """Navigate back to the previous page in the browser history."""
-
- name: str = "previous_webpage"
- description: str = "Navigate back to the previous page in the browser history"
- args_schema: Type[BaseModel] = NavigateBackToolInput
-
- def _run(self, run_manager: Optional[CallbackManagerForToolRun] = None) -> str:
- """Use the tool."""
- if self.sync_browser is None:
- raise ValueError(f"Synchronous browser not provided to {self.name}")
- page = get_current_page(self.sync_browser)
- response = page.go_back()
-
- if response:
- return (
- f"Navigated back to the previous page with URL '{response.url}'."
- f" Status code {response.status}"
- )
- else:
- return "Unable to navigate back; no previous page in the history"
-
- async def _arun(
- self,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- if self.async_browser is None:
- raise ValueError(f"Asynchronous browser not provided to {self.name}")
- page = await aget_current_page(self.async_browser)
- response = await page.go_back()
-
- if response:
- return (
- f"Navigated back to the previous page with URL '{response.url}'."
- f" Status code {response.status}"
- )
- else:
- return "Unable to navigate back; no previous page in the history"
diff --git a/libs/community/langchain_community/tools/playwright/utils.py b/libs/community/langchain_community/tools/playwright/utils.py
deleted file mode 100644
index 9373873662..0000000000
--- a/libs/community/langchain_community/tools/playwright/utils.py
+++ /dev/null
@@ -1,105 +0,0 @@
-"""Utilities for the Playwright browser tools."""
-
-from __future__ import annotations
-
-import asyncio
-from typing import TYPE_CHECKING, Any, Coroutine, List, Optional, TypeVar
-
-if TYPE_CHECKING:
- from playwright.async_api import Browser as AsyncBrowser
- from playwright.async_api import Page as AsyncPage
- from playwright.sync_api import Browser as SyncBrowser
- from playwright.sync_api import Page as SyncPage
-
-
-async def aget_current_page(browser: AsyncBrowser) -> AsyncPage:
- """
- Asynchronously get the current page of the browser.
-
- Args:
- browser: The browser (AsyncBrowser) to get the current page from.
-
- Returns:
- AsyncPage: The current page.
- """
- if not browser.contexts:
- context = await browser.new_context()
- return await context.new_page()
- context = browser.contexts[0] # Assuming you're using the default browser context
- if not context.pages:
- return await context.new_page()
- # Assuming the last page in the list is the active one
- return context.pages[-1]
-
-
-def get_current_page(browser: SyncBrowser) -> SyncPage:
- """
- Get the current page of the browser.
- Args:
- browser: The browser to get the current page from.
-
- Returns:
- SyncPage: The current page.
- """
- if not browser.contexts:
- context = browser.new_context()
- return context.new_page()
- context = browser.contexts[0] # Assuming you're using the default browser context
- if not context.pages:
- return context.new_page()
- # Assuming the last page in the list is the active one
- return context.pages[-1]
-
-
-def create_async_playwright_browser(
- headless: bool = True, args: Optional[List[str]] = None
-) -> AsyncBrowser:
- """
- Create an async playwright browser.
-
- Args:
- headless: Whether to run the browser in headless mode. Defaults to True.
- args: arguments to pass to browser.chromium.launch
-
- Returns:
- AsyncBrowser: The playwright browser.
- """
- from playwright.async_api import async_playwright
-
- browser = run_async(async_playwright().start())
- return run_async(browser.chromium.launch(headless=headless, args=args))
-
-
-def create_sync_playwright_browser(
- headless: bool = True, args: Optional[List[str]] = None
-) -> SyncBrowser:
- """
- Create a playwright browser.
-
- Args:
- headless: Whether to run the browser in headless mode. Defaults to True.
- args: arguments to pass to browser.chromium.launch
-
- Returns:
- SyncBrowser: The playwright browser.
- """
- from playwright.sync_api import sync_playwright
-
- browser = sync_playwright().start()
- return browser.chromium.launch(headless=headless, args=args)
-
-
-T = TypeVar("T")
-
-
-def run_async(coro: Coroutine[Any, Any, T]) -> T:
- """Run an async coroutine.
-
- Args:
- coro: The coroutine to run. Coroutine[Any, Any, T]
-
- Returns:
- T: The result of the coroutine.
- """
- event_loop = asyncio.get_event_loop()
- return event_loop.run_until_complete(coro)
diff --git a/libs/community/langchain_community/tools/plugin.py b/libs/community/langchain_community/tools/plugin.py
deleted file mode 100644
index 102451e72d..0000000000
--- a/libs/community/langchain_community/tools/plugin.py
+++ /dev/null
@@ -1,110 +0,0 @@
-from __future__ import annotations
-
-import json
-from typing import Optional, Type
-
-import requests
-import yaml
-from langchain_core.callbacks import (
- AsyncCallbackManagerForToolRun,
- CallbackManagerForToolRun,
-)
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel
-
-
-class ApiConfig(BaseModel):
- """API Configuration."""
-
- type: str
- url: str
- has_user_authentication: Optional[bool] = False
-
-
-class AIPlugin(BaseModel):
- """AI Plugin Definition."""
-
- schema_version: str
- name_for_model: str
- name_for_human: str
- description_for_model: str
- description_for_human: str
- auth: Optional[dict] = None
- api: ApiConfig
- logo_url: Optional[str]
- contact_email: Optional[str]
- legal_info_url: Optional[str]
-
- @classmethod
- def from_url(cls, url: str) -> AIPlugin:
- """Instantiate AIPlugin from a URL."""
- response = requests.get(url).json()
- return cls(**response)
-
-
-def marshal_spec(txt: str) -> dict:
- """Convert the yaml or json serialized spec to a dict.
-
- Args:
- txt: The yaml or json serialized spec.
-
- Returns:
- dict: The spec as a dict.
- """
- try:
- return json.loads(txt)
- except json.JSONDecodeError:
- return yaml.safe_load(txt)
-
-
-class AIPluginToolSchema(BaseModel):
- """Schema for AIPluginTool."""
-
- tool_input: Optional[str] = ""
-
-
-class AIPluginTool(BaseTool):
- """Tool for getting the OpenAPI spec for an AI Plugin."""
-
- plugin: AIPlugin
- api_spec: str
- args_schema: Type[AIPluginToolSchema] = AIPluginToolSchema
-
- @classmethod
- def from_plugin_url(cls, url: str) -> AIPluginTool:
- plugin = AIPlugin.from_url(url)
- description = (
- f"Call this tool to get the OpenAPI spec (and usage guide) "
- f"for interacting with the {plugin.name_for_human} API. "
- f"You should only call this ONCE! What is the "
- f"{plugin.name_for_human} API useful for? "
- ) + plugin.description_for_human
- open_api_spec_str = requests.get(plugin.api.url).text
- open_api_spec = marshal_spec(open_api_spec_str)
- api_spec = (
- f"Usage Guide: {plugin.description_for_model}\n\n"
- f"OpenAPI Spec: {open_api_spec}"
- )
-
- return cls(
- name=plugin.name_for_model,
- description=description,
- plugin=plugin,
- api_spec=api_spec,
- )
-
- def _run(
- self,
- tool_input: Optional[str] = "",
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool."""
- return self.api_spec
-
- async def _arun(
- self,
- tool_input: Optional[str] = None,
- run_manager: Optional[AsyncCallbackManagerForToolRun] = None,
- ) -> str:
- """Use the tool asynchronously."""
- return self.api_spec
diff --git a/libs/community/langchain_community/tools/polygon/__init__.py b/libs/community/langchain_community/tools/polygon/__init__.py
deleted file mode 100644
index 87a9c1c135..0000000000
--- a/libs/community/langchain_community/tools/polygon/__init__.py
+++ /dev/null
@@ -1,13 +0,0 @@
-"""Polygon IO tools."""
-
-from langchain_community.tools.polygon.aggregates import PolygonAggregates
-from langchain_community.tools.polygon.financials import PolygonFinancials
-from langchain_community.tools.polygon.last_quote import PolygonLastQuote
-from langchain_community.tools.polygon.ticker_news import PolygonTickerNews
-
-__all__ = [
- "PolygonAggregates",
- "PolygonFinancials",
- "PolygonLastQuote",
- "PolygonTickerNews",
-]
diff --git a/libs/community/langchain_community/tools/polygon/aggregates.py b/libs/community/langchain_community/tools/polygon/aggregates.py
deleted file mode 100644
index 26cb62d467..0000000000
--- a/libs/community/langchain_community/tools/polygon/aggregates.py
+++ /dev/null
@@ -1,77 +0,0 @@
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel, Field
-
-from langchain_community.utilities.polygon import PolygonAPIWrapper
-
-
-class PolygonAggregatesSchema(BaseModel):
- """Input for PolygonAggregates."""
-
- ticker: str = Field(
- description="The ticker symbol to fetch aggregates for.",
- )
- timespan: str = Field(
- description="The size of the time window. "
- "Possible values are: "
- "second, minute, hour, day, week, month, quarter, year. "
- "Default is 'day'",
- )
- timespan_multiplier: int = Field(
- description="The number of timespans to aggregate. "
- "For example, if timespan is 'day' and "
- "timespan_multiplier is 1, the result will be daily bars. "
- "If timespan is 'day' and timespan_multiplier is 5, "
- "the result will be weekly bars. "
- "Default is 1.",
- )
- from_date: str = Field(
- description="The start of the aggregate time window. "
- "Either a date with the format YYYY-MM-DD or "
- "a millisecond timestamp.",
- )
- to_date: str = Field(
- description="The end of the aggregate time window. "
- "Either a date with the format YYYY-MM-DD or "
- "a millisecond timestamp.",
- )
-
-
-class PolygonAggregates(BaseTool):
- """
- Tool that gets aggregate bars (stock prices) over a
- given date range for a given ticker from Polygon.
- """
-
- mode: str = "get_aggregates"
- name: str = "polygon_aggregates"
- description: str = (
- "A wrapper around Polygon's Aggregates API. "
- "This tool is useful for fetching aggregate bars (stock prices) for a ticker. "
- "Input should be the ticker, date range, timespan, and timespan multiplier"
- " that you want to get the aggregate bars for."
- )
- args_schema: Type[PolygonAggregatesSchema] = PolygonAggregatesSchema
-
- api_wrapper: PolygonAPIWrapper
-
- def _run(
- self,
- ticker: str,
- timespan: str,
- timespan_multiplier: int,
- from_date: str,
- to_date: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Polygon API tool."""
- return self.api_wrapper.run(
- mode=self.mode,
- ticker=ticker,
- timespan=timespan,
- timespan_multiplier=timespan_multiplier,
- from_date=from_date,
- to_date=to_date,
- )
diff --git a/libs/community/langchain_community/tools/polygon/financials.py b/libs/community/langchain_community/tools/polygon/financials.py
deleted file mode 100644
index 8400e7498b..0000000000
--- a/libs/community/langchain_community/tools/polygon/financials.py
+++ /dev/null
@@ -1,38 +0,0 @@
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel
-
-from langchain_community.utilities.polygon import PolygonAPIWrapper
-
-
-class Inputs(BaseModel):
- """Inputs for Polygon's Financials API"""
-
- query: str
-
-
-class PolygonFinancials(BaseTool):
- """Tool that gets the financials of a ticker from Polygon"""
-
- mode: str = "get_financials"
- name: str = "polygon_financials"
- description: str = (
- "A wrapper around Polygon's Stock Financials API. "
- "This tool is useful for fetching fundamental financials from "
- "balance sheets, income statements, and cash flow statements "
- "for a stock ticker. The input should be the ticker that you want "
- "to get the latest fundamental financial data for."
- )
- args_schema: Type[BaseModel] = Inputs
-
- api_wrapper: PolygonAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Polygon API tool."""
- return self.api_wrapper.run(self.mode, ticker=query)
diff --git a/libs/community/langchain_community/tools/polygon/last_quote.py b/libs/community/langchain_community/tools/polygon/last_quote.py
deleted file mode 100644
index 76c768113b..0000000000
--- a/libs/community/langchain_community/tools/polygon/last_quote.py
+++ /dev/null
@@ -1,36 +0,0 @@
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel
-
-from langchain_community.utilities.polygon import PolygonAPIWrapper
-
-
-class Inputs(BaseModel):
- """Inputs for Polygon's Last Quote API"""
-
- query: str
-
-
-class PolygonLastQuote(BaseTool):
- """Tool that gets the last quote of a ticker from Polygon"""
-
- mode: str = "get_last_quote"
- name: str = "polygon_last_quote"
- description: str = (
- "A wrapper around Polygon's Last Quote API. "
- "This tool is useful for fetching the latest price of a stock. "
- "Input should be the ticker that you want to query the last price quote for."
- )
- args_schema: Type[BaseModel] = Inputs
-
- api_wrapper: PolygonAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Polygon API tool."""
- return self.api_wrapper.run(self.mode, ticker=query)
diff --git a/libs/community/langchain_community/tools/polygon/ticker_news.py b/libs/community/langchain_community/tools/polygon/ticker_news.py
deleted file mode 100644
index d4c4a2017a..0000000000
--- a/libs/community/langchain_community/tools/polygon/ticker_news.py
+++ /dev/null
@@ -1,36 +0,0 @@
-from typing import Optional, Type
-
-from langchain_core.callbacks import CallbackManagerForToolRun
-from langchain_core.tools import BaseTool
-from pydantic import BaseModel
-
-from langchain_community.utilities.polygon import PolygonAPIWrapper
-
-
-class Inputs(BaseModel):
- """Inputs for Polygon's Ticker News API"""
-
- query: str
-
-
-class PolygonTickerNews(BaseTool):
- """Tool that gets the latest news for a given ticker from Polygon"""
-
- mode: str = "get_ticker_news"
- name: str = "polygon_ticker_news"
- description: str = (
- "A wrapper around Polygon's Ticker News API. "
- "This tool is useful for fetching the latest news for a stock. "
- "Input should be the ticker that you want to get the latest news for."
- )
- args_schema: Type[BaseModel] = Inputs
-
- api_wrapper: PolygonAPIWrapper
-
- def _run(
- self,
- query: str,
- run_manager: Optional[CallbackManagerForToolRun] = None,
- ) -> str:
- """Use the Polygon API tool."""
- return self.api_wrapper.run(self.mode, ticker=query)
diff --git a/libs/community/langchain_community/tools/powerbi/__init__.py b/libs/community/langchain_community/tools/powerbi/__init__.py
deleted file mode 100644
index 3ecc25a12f..0000000000
--- a/libs/community/langchain_community/tools/powerbi/__init__.py
+++ /dev/null
@@ -1 +0,0 @@
-"""Tools for interacting with a PowerBI dataset."""
diff --git a/libs/community/langchain_community/tools/powerbi/prompt.py b/libs/community/langchain_community/tools/powerbi/prompt.py
deleted file mode 100644
index caf32756ac..0000000000
--- a/libs/community/langchain_community/tools/powerbi/prompt.py
+++ /dev/null
@@ -1,70 +0,0 @@
-# flake8: noqa
-QUESTION_TO_QUERY_BASE = """
-Answer the question below with a DAX query that can be sent to Power BI. DAX queries have a simple syntax comprised of just one required keyword, EVALUATE, and several optional keywords: ORDER BY, START AT, DEFINE, MEASURE, VAR, TABLE, and COLUMN. Each keyword defines a statement used for the duration of the query. Any time < or > are used in the text below it means that those values need to be replaced by table, columns or other things. If the question is not something you can answer with a DAX query, reply with "I cannot answer this" and the question will be escalated to a human.
-
-Some DAX functions return a table instead of a scalar, and must be wrapped in a function that evaluates the table and returns a scalar; unless the table is a single column, single row table, then it is treated as a scalar value. Most DAX functions require one or more arguments, which can include tables, columns, expressions, and values. However, some functions, such as PI, do not require any arguments, but always require parentheses to indicate the null argument. For example, you must always type PI(), not PI. You can also nest functions within other functions.
-
-Some commonly used functions are:
-EVALUATE
- At the most basic level, a DAX query is an EVALUATE statement containing a table expression. At least one EVALUATE statement is required, however, a query can contain any number of EVALUATE statements.
-EVALUATE
ORDER BY ASC or DESC - The optional ORDER BY keyword defines one or more expressions used to sort query results. Any expression that can be evaluated for each row of the result is valid.
-EVALUATE
ORDER BY ASC or DESC START AT or - The optional START AT keyword is used inside an ORDER BY clause. It defines the value at which the query results begin.
-DEFINE MEASURE | VAR; EVALUATE
- The optional DEFINE keyword introduces one or more calculated entity definitions that exist only for the duration of the query. Definitions precede the EVALUATE statement and are valid for all EVALUATE statements in the query. Definitions can be variables, measures, tables1, and columns1. Definitions can reference other definitions that appear before or after the current definition. At least one definition is required if the DEFINE keyword is included in a query.
-MEASURE
[] = - Introduces a measure definition in a DEFINE statement of a DAX query.
-VAR = - Stores the result of an expression as a named variable, which can then be passed as an argument to other measure expressions. Once resultant values have been calculated for a variable expression, those values do not change, even if the variable is referenced in another expression.
-
-FILTER(
,) - Returns a table that represents a subset of another table or expression, where is a Boolean expression that is to be evaluated for each row of the table. For example, [Amount] > 0 or [Region] = "France"
-ROW(, ) - Returns a table with a single row containing values that result from the expressions given to each column.
-TOPN(,
, , ) - Returns a table with the top n rows from the specified table, sorted by the specified expression, in the order specified by 0 for descending, 1 for ascending, the default is 0. Multiple OrderBy_Expressions and Order pairs can be given, separated by a comma.
-DISTINCT() - Returns a one-column table that contains the distinct values from the specified column. In other words, duplicate values are removed and only unique values are returned. This function cannot be used to Return values into a cell or column on a worksheet; rather, you nest the DISTINCT function within a formula, to get a list of distinct values that can be passed to another function and then counted, summed, or used for other operations.
-DISTINCT(
) - Returns a table by removing duplicate rows from another table or expression.
-
-Aggregation functions, names with a A in it, handle booleans and empty strings in appropriate ways, while the same function without A only uses the numeric values in a column. Functions names with an X in it can include a expression as an argument, this will be evaluated for each row in the table and the result will be used in the regular function calculation, these are the functions:
-COUNT(), COUNTA(), COUNTX(
,), COUNTAX(
,), COUNTROWS([
]), COUNTBLANK(), DISTINCTCOUNT(), DISTINCTCOUNTNOBLANK () - these are all variations of count functions.
-AVERAGE(), AVERAGEA(), AVERAGEX(
,) - these are all variations of average functions.
-MAX(), MAXA(), MAXX(
,) - these are all variations of max functions.
-MIN(), MINA(), MINX(
,) - these are all variations of min functions.
-PRODUCT(), PRODUCTX(
,) - these are all variations of product functions.
-SUM(), SUMX(
,) - these are all variations of sum functions.
-
-Date and time functions:
-DATE(year, month, day) - Returns a date value that represents the specified year, month, and day.
-DATEDIFF(date1, date2, ) - Returns the difference between two date values, in the specified interval, that can be SECOND, MINUTE, HOUR, DAY, WEEK, MONTH, QUARTER, YEAR.
-DATEVALUE() - Returns a date value that represents the specified date.
-YEAR(), QUARTER(), MONTH(), DAY(), HOUR(), MINUTE(), SECOND() - Returns the part of the date for the specified date.
-
-Finally, make sure to escape double quotes with a single backslash, and make sure that only table names have single quotes around them, while names of measures or the values of columns that you want to compare against are in escaped double quotes. Newlines are not necessary and can be skipped. The queries are serialized as json and so will have to fit be compliant with json syntax. Sometimes you will get a question, a DAX query and a error, in that case you need to rewrite the DAX query to get the correct answer.
-
-The following tables exist: {tables}
-
-and the schema's for some are given here:
-{schemas}
-
-Examples:
-{examples}
-"""
-
-USER_INPUT = """
-Question: {tool_input}
-DAX:
-"""
-
-SINGLE_QUESTION_TO_QUERY = f"{QUESTION_TO_QUERY_BASE}{USER_INPUT}"
-
-DEFAULT_FEWSHOT_EXAMPLES = """
-Question: How many rows are in the table
?
-DAX: EVALUATE ROW(\"Number of rows\", COUNTROWS(
))
-----
-Question: How many rows are in the table
where is not empty?
-DAX: EVALUATE ROW(\"Number of rows\", COUNTROWS(FILTER(