mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
chore(core): fix some any generics (#34545)
Co-authored-by: Mason Daugherty <github@mdrxy.com>
This commit is contained in:
1 parent
3eee4002d9
commit
a063ec26dd
100 files changed
+689
-531
No files matched your search
@@ -331,7 +331,7 @@ def deprecated(
|
||||
if not _obj_type:
|
||||
_obj_type = "attribute"
|
||||
wrapped = None
|
||||
_name = _name or cast("type | Callable", obj.fget).__qualname__
|
||||
_name = _name or cast("type", obj.fget).__qualname__
|
||||
old_doc = obj.__doc__
|
||||
|
||||
class _DeprecatedProperty(property):
|
||||
@@ -386,7 +386,7 @@ def deprecated(
|
||||
return cast("T", prop)
|
||||
|
||||
else:
|
||||
_name = _name or cast("type | Callable", obj).__qualname__
|
||||
_name = _name or cast("type", obj).__qualname__
|
||||
if not _obj_type:
|
||||
# edge case: when a function is within another function
|
||||
# within a test, this will call it a "method" not a "function"
|
||||
|
||||
@@ -51,7 +51,7 @@ class AgentAction(Serializable):
|
||||
tool: str
|
||||
"""The name of the `Tool` to execute."""
|
||||
|
||||
tool_input: str | dict
|
||||
tool_input: str | dict[Any, Any]
|
||||
"""The input to pass in to the `Tool`."""
|
||||
|
||||
log: str
|
||||
@@ -68,7 +68,9 @@ class AgentAction(Serializable):
|
||||
type: Literal["AgentAction"] = "AgentAction"
|
||||
|
||||
# Override init to support instantiation by position for backward compat.
|
||||
def __init__(self, tool: str, tool_input: str | dict, log: str, **kwargs: Any):
|
||||
def __init__(
|
||||
self, tool: str, tool_input: str | dict[Any, Any], log: str, **kwargs: Any
|
||||
):
|
||||
"""Create an `AgentAction`.
|
||||
|
||||
Args:
|
||||
@@ -149,7 +151,7 @@ class AgentFinish(Serializable):
|
||||
Agents return an `AgentFinish` when they have reached a stopping condition.
|
||||
"""
|
||||
|
||||
return_values: dict
|
||||
return_values: dict[Any, Any]
|
||||
"""Dictionary of return values."""
|
||||
|
||||
log: str
|
||||
@@ -164,7 +166,7 @@ class AgentFinish(Serializable):
|
||||
"""
|
||||
type: Literal["AgentFinish"] = "AgentFinish"
|
||||
|
||||
def __init__(self, return_values: dict, log: str, **kwargs: Any):
|
||||
def __init__(self, return_values: dict[Any, Any], log: str, **kwargs: Any):
|
||||
"""Override init to support instantiation by position for backward compat."""
|
||||
super().__init__(return_values=return_values, log=log, **kwargs)
|
||||
|
||||
|
||||
@@ -214,7 +214,7 @@ async def atrace_as_chain_group(
|
||||
await run_manager.on_chain_end({})
|
||||
|
||||
|
||||
Func = TypeVar("Func", bound=Callable)
|
||||
Func = TypeVar("Func", bound=Callable[..., Any])
|
||||
|
||||
|
||||
def shielded(func: Func) -> Func:
|
||||
@@ -327,9 +327,7 @@ def handle_event(
|
||||
# running coroutine, which we cannot interrupt to run this one.
|
||||
# The solution is to run the synchronous function on the globally shared
|
||||
# thread pool executor to avoid blocking the main event loop.
|
||||
_executor().submit(
|
||||
cast("Callable", copy_context().run), _run_coros, coros
|
||||
).result()
|
||||
_executor().submit(copy_context().run, _run_coros, coros).result()
|
||||
else:
|
||||
# If there's no running loop, we can run the coroutines directly.
|
||||
_run_coros(coros)
|
||||
@@ -381,10 +379,7 @@ async def _ahandle_event_for_handler(
|
||||
else:
|
||||
await asyncio.get_event_loop().run_in_executor(
|
||||
None,
|
||||
cast(
|
||||
"Callable",
|
||||
functools.partial(copy_context().run, event, *args, **kwargs),
|
||||
),
|
||||
functools.partial(copy_context().run, event, *args, **kwargs),
|
||||
)
|
||||
except NotImplementedError as e:
|
||||
if event_name == "on_chat_model_start":
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""**Chat Sessions** are a collection of messages and function calls."""
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import TypedDict
|
||||
from typing import Any, TypedDict
|
||||
|
||||
from langchain_core.messages import BaseMessage
|
||||
|
||||
@@ -15,5 +15,5 @@ class ChatSession(TypedDict, total=False):
|
||||
messages: Sequence[BaseMessage]
|
||||
"""A sequence of the LangChain chat messages loaded from the source."""
|
||||
|
||||
functions: Sequence[dict]
|
||||
functions: Sequence[dict[str, Any]]
|
||||
"""A sequence of the function calling specs for the messages."""
|
||||
@@ -48,7 +48,7 @@ class LangSmithLoader(BaseLoader):
|
||||
inline_s3_urls: bool = True,
|
||||
offset: int = 0,
|
||||
limit: int | None = None,
|
||||
metadata: dict | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
filter: str | None = None, # noqa: A002
|
||||
content_key: str = "",
|
||||
format_content: Callable[..., str] | None = None,
|
||||
|
||||
@@ -52,7 +52,7 @@ class BaseMedia(Serializable):
|
||||
as a UUID, but this will not be enforced.
|
||||
"""
|
||||
|
||||
metadata: dict = Field(default_factory=dict)
|
||||
metadata: dict[Any, Any] = Field(default_factory=dict)
|
||||
"""Arbitrary metadata associated with the content."""
|
||||
|
||||
|
||||
@@ -218,7 +218,7 @@ class Blob(BaseMedia):
|
||||
encoding: str = "utf-8",
|
||||
mime_type: str | None = None,
|
||||
guess_type: bool = True,
|
||||
metadata: dict | None = None,
|
||||
metadata: dict[Any, Any] | None = None,
|
||||
) -> Blob:
|
||||
"""Load the blob from a path like object.
|
||||
|
||||
@@ -255,7 +255,7 @@ class Blob(BaseMedia):
|
||||
encoding: str = "utf-8",
|
||||
mime_type: str | None = None,
|
||||
path: str | None = None,
|
||||
metadata: dict | None = None,
|
||||
metadata: dict[Any, Any] | None = None,
|
||||
) -> Blob:
|
||||
"""Initialize the `Blob` from in-memory data.
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ from langchain_core import __version__
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def get_runtime_environment() -> dict:
|
||||
def get_runtime_environment() -> dict[str, str]:
|
||||
"""Get information about the LangChain runtime environment.
|
||||
|
||||
Returns:
|
||||
|
||||
@@ -34,7 +34,7 @@ class BaseExampleSelector(ABC):
|
||||
return await run_in_executor(None, self.add_example, example)
|
||||
|
||||
@abstractmethod
|
||||
def select_examples(self, input_variables: dict[str, str]) -> list[dict]:
|
||||
def select_examples(self, input_variables: dict[str, str]) -> list[dict[str, Any]]:
|
||||
"""Select which examples to use based on the inputs.
|
||||
|
||||
Args:
|
||||
@@ -45,7 +45,9 @@ class BaseExampleSelector(ABC):
|
||||
A list of examples.
|
||||
"""
|
||||
|
||||
async def aselect_examples(self, input_variables: dict[str, str]) -> list[dict]:
|
||||
async def aselect_examples(
|
||||
self, input_variables: dict[str, str]
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Async select which examples to use based on the inputs.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
import re
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from typing_extensions import Self
|
||||
@@ -48,7 +49,7 @@ class LengthBasedExampleSelector(BaseExampleSelector, BaseModel):
|
||||
```
|
||||
"""
|
||||
|
||||
examples: list[dict]
|
||||
examples: list[dict[str, Any]]
|
||||
"""A list of the examples that the prompt template expects."""
|
||||
|
||||
example_prompt: PromptTemplate
|
||||
@@ -92,12 +93,12 @@ class LengthBasedExampleSelector(BaseExampleSelector, BaseModel):
|
||||
self.example_text_lengths = [self.get_text_length(eg) for eg in string_examples]
|
||||
return self
|
||||
|
||||
def select_examples(self, input_variables: dict[str, str]) -> list[dict]:
|
||||
def select_examples(self, input_variables: dict[str, str]) -> list[dict[str, Any]]:
|
||||
"""Select which examples to use based on the input lengths.
|
||||
|
||||
Args:
|
||||
input_variables: A dictionary with keys as input variables
|
||||
and values as their values.
|
||||
and values as their values.
|
||||
|
||||
Returns:
|
||||
A list of examples to include in the prompt.
|
||||
@@ -115,12 +116,14 @@ class LengthBasedExampleSelector(BaseExampleSelector, BaseModel):
|
||||
i += 1
|
||||
return examples
|
||||
|
||||
async def aselect_examples(self, input_variables: dict[str, str]) -> list[dict]:
|
||||
async def aselect_examples(
|
||||
self, input_variables: dict[str, str]
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Async select which examples to use based on the input lengths.
|
||||
|
||||
Args:
|
||||
input_variables: A dictionary with keys as input variables
|
||||
and values as their values.
|
||||
and values as their values.
|
||||
|
||||
Returns:
|
||||
A list of examples to include in the prompt.
|
||||
|
||||
@@ -15,7 +15,7 @@ if TYPE_CHECKING:
|
||||
from langchain_core.embeddings import Embeddings
|
||||
|
||||
|
||||
def sorted_values(values: dict[str, str]) -> list[Any]:
|
||||
def sorted_values(values: dict[str, str]) -> list[str]:
|
||||
"""Return a list of values in dict sorted by key.
|
||||
|
||||
Args:
|
||||
@@ -33,13 +33,17 @@ class _VectorStoreExampleSelector(BaseExampleSelector, BaseModel, ABC):
|
||||
|
||||
vectorstore: VectorStore
|
||||
"""VectorStore that contains information about examples."""
|
||||
|
||||
k: int = 4
|
||||
"""Number of examples to select."""
|
||||
|
||||
example_keys: list[str] | None = None
|
||||
"""Optional keys to filter examples to."""
|
||||
|
||||
input_keys: list[str] | None = None
|
||||
"""Optional keys to filter input to. If provided, the search is based on
|
||||
the input variables instead of all variables."""
|
||||
|
||||
vectorstore_kwargs: dict[str, Any] | None = None
|
||||
"""Extra arguments passed to similarity_search function of the `VectorStore`."""
|
||||
|
||||
@@ -54,7 +58,7 @@ class _VectorStoreExampleSelector(BaseExampleSelector, BaseModel, ABC):
|
||||
return " ".join(sorted_values({key: example[key] for key in input_keys}))
|
||||
return " ".join(sorted_values(example))
|
||||
|
||||
def _documents_to_examples(self, documents: list[Document]) -> list[dict]:
|
||||
def _documents_to_examples(self, documents: list[Document]) -> list[dict[str, Any]]:
|
||||
# Get the examples from the metadata.
|
||||
# This assumes that examples are stored in metadata.
|
||||
examples = [dict(e.metadata) for e in documents]
|
||||
@@ -97,7 +101,7 @@ class _VectorStoreExampleSelector(BaseExampleSelector, BaseModel, ABC):
|
||||
class SemanticSimilarityExampleSelector(_VectorStoreExampleSelector):
|
||||
"""Select examples based on semantic similarity."""
|
||||
|
||||
def select_examples(self, input_variables: dict[str, str]) -> list[dict]:
|
||||
def select_examples(self, input_variables: dict[str, str]) -> list[dict[str, Any]]:
|
||||
"""Select examples based on semantic similarity.
|
||||
|
||||
Args:
|
||||
@@ -115,7 +119,9 @@ class SemanticSimilarityExampleSelector(_VectorStoreExampleSelector):
|
||||
)
|
||||
return self._documents_to_examples(example_docs)
|
||||
|
||||
async def aselect_examples(self, input_variables: dict[str, str]) -> list[dict]:
|
||||
async def aselect_examples(
|
||||
self, input_variables: dict[str, str]
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Asynchronously select examples based on semantic similarity.
|
||||
|
||||
Args:
|
||||
@@ -136,14 +142,14 @@ class SemanticSimilarityExampleSelector(_VectorStoreExampleSelector):
|
||||
@classmethod
|
||||
def from_examples(
|
||||
cls,
|
||||
examples: list[dict],
|
||||
examples: list[dict[str, str]],
|
||||
embeddings: Embeddings,
|
||||
vectorstore_cls: type[VectorStore],
|
||||
k: int = 4,
|
||||
input_keys: list[str] | None = None,
|
||||
*,
|
||||
example_keys: list[str] | None = None,
|
||||
vectorstore_kwargs: dict | None = None,
|
||||
vectorstore_kwargs: dict[str, Any] | None = None,
|
||||
**vectorstore_cls_kwargs: Any,
|
||||
) -> SemanticSimilarityExampleSelector:
|
||||
"""Create k-shot example selector using example list and embeddings.
|
||||
@@ -180,14 +186,14 @@ class SemanticSimilarityExampleSelector(_VectorStoreExampleSelector):
|
||||
@classmethod
|
||||
async def afrom_examples(
|
||||
cls,
|
||||
examples: list[dict],
|
||||
examples: list[dict[str, str]],
|
||||
embeddings: Embeddings,
|
||||
vectorstore_cls: type[VectorStore],
|
||||
k: int = 4,
|
||||
input_keys: list[str] | None = None,
|
||||
*,
|
||||
example_keys: list[str] | None = None,
|
||||
vectorstore_kwargs: dict | None = None,
|
||||
vectorstore_kwargs: dict[str, Any] | None = None,
|
||||
**vectorstore_cls_kwargs: Any,
|
||||
) -> SemanticSimilarityExampleSelector:
|
||||
"""Async create k-shot example selector using example list and embeddings.
|
||||
@@ -232,7 +238,7 @@ class MaxMarginalRelevanceExampleSelector(_VectorStoreExampleSelector):
|
||||
fetch_k: int = 20
|
||||
"""Number of examples to fetch to rerank."""
|
||||
|
||||
def select_examples(self, input_variables: dict[str, str]) -> list[dict]:
|
||||
def select_examples(self, input_variables: dict[str, str]) -> list[dict[str, Any]]:
|
||||
"""Select examples based on Max Marginal Relevance.
|
||||
|
||||
Args:
|
||||
@@ -248,7 +254,9 @@ class MaxMarginalRelevanceExampleSelector(_VectorStoreExampleSelector):
|
||||
)
|
||||
return self._documents_to_examples(example_docs)
|
||||
|
||||
async def aselect_examples(self, input_variables: dict[str, str]) -> list[dict]:
|
||||
async def aselect_examples(
|
||||
self, input_variables: dict[str, str]
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Asynchronously select examples based on Max Marginal Relevance.
|
||||
|
||||
Args:
|
||||
@@ -267,14 +275,14 @@ class MaxMarginalRelevanceExampleSelector(_VectorStoreExampleSelector):
|
||||
@classmethod
|
||||
def from_examples(
|
||||
cls,
|
||||
examples: list[dict],
|
||||
examples: list[dict[str, str]],
|
||||
embeddings: Embeddings,
|
||||
vectorstore_cls: type[VectorStore],
|
||||
k: int = 4,
|
||||
input_keys: list[str] | None = None,
|
||||
fetch_k: int = 20,
|
||||
example_keys: list[str] | None = None,
|
||||
vectorstore_kwargs: dict | None = None,
|
||||
vectorstore_kwargs: dict[str, Any] | None = None,
|
||||
**vectorstore_cls_kwargs: Any,
|
||||
) -> MaxMarginalRelevanceExampleSelector:
|
||||
"""Create k-shot example selector using example list and embeddings.
|
||||
@@ -313,7 +321,7 @@ class MaxMarginalRelevanceExampleSelector(_VectorStoreExampleSelector):
|
||||
@classmethod
|
||||
async def afrom_examples(
|
||||
cls,
|
||||
examples: list[dict],
|
||||
examples: list[dict[str, str]],
|
||||
embeddings: Embeddings,
|
||||
vectorstore_cls: type[VectorStore],
|
||||
*,
|
||||
@@ -321,7 +329,7 @@ class MaxMarginalRelevanceExampleSelector(_VectorStoreExampleSelector):
|
||||
input_keys: list[str] | None = None,
|
||||
fetch_k: int = 20,
|
||||
example_keys: list[str] | None = None,
|
||||
vectorstore_kwargs: dict | None = None,
|
||||
vectorstore_kwargs: dict[str, Any] | None = None,
|
||||
**vectorstore_cls_kwargs: Any,
|
||||
) -> MaxMarginalRelevanceExampleSelector:
|
||||
"""Create k-shot example selector using example list and embeddings.
|
||||
|
||||
@@ -31,7 +31,7 @@ def _filter_invocation_params_for_tracing(params: dict[str, Any]) -> dict[str, A
|
||||
|
||||
|
||||
def is_openai_data_block(
|
||||
block: dict, filter_: Literal["image", "audio", "file"] | None = None
|
||||
block: dict[str, Any], filter_: Literal["image", "audio", "file"] | None = None
|
||||
) -> bool:
|
||||
"""Check whether a block contains multimodal data in OpenAI Chat Completions format.
|
||||
|
||||
@@ -318,7 +318,7 @@ def _ensure_message_copy(message: T, formatted_message: T) -> T:
|
||||
|
||||
|
||||
def _update_content_block(
|
||||
formatted_message: "BaseMessage", idx: int, new_block: ContentBlock | dict
|
||||
formatted_message: "BaseMessage", idx: int, new_block: ContentBlock | dict[str, Any]
|
||||
) -> None:
|
||||
"""Update a content block at the given index, handling type issues."""
|
||||
# Type ignore needed because:
|
||||
|
||||
@@ -294,8 +294,8 @@ class BaseLanguageModel(
|
||||
"""
|
||||
|
||||
def with_structured_output(
|
||||
self, schema: dict | type, **kwargs: Any
|
||||
) -> Runnable[LanguageModelInput, dict | BaseModel]:
|
||||
self, schema: dict[str, Any] | type, **kwargs: Any
|
||||
) -> Runnable[LanguageModelInput, dict[str, Any] | BaseModel]:
|
||||
"""Not implemented on this class."""
|
||||
# Implement this on child class if there is a way of steering the model to
|
||||
# generate responses that match a given schema.
|
||||
@@ -356,7 +356,7 @@ class BaseLanguageModel(
|
||||
def get_num_tokens_from_messages(
|
||||
self,
|
||||
messages: list[BaseMessage],
|
||||
tools: Sequence | None = None,
|
||||
tools: Sequence[Any] | None = None,
|
||||
) -> int:
|
||||
"""Get the number of tokens in the messages.
|
||||
|
||||
|
||||
@@ -67,6 +67,7 @@ from langchain_core.messages.block_translators.openai import (
|
||||
)
|
||||
from langchain_core.output_parsers.openai_tools import (
|
||||
JsonOutputKeyToolsParser,
|
||||
JsonOutputToolsParser,
|
||||
PydanticToolsParser,
|
||||
)
|
||||
from langchain_core.outputs import (
|
||||
@@ -99,7 +100,6 @@ if TYPE_CHECKING:
|
||||
|
||||
from langchain_protocol.protocol import MessagesData
|
||||
|
||||
from langchain_core.output_parsers.base import OutputParserLike
|
||||
from langchain_core.runnables import Runnable, RunnableConfig
|
||||
from langchain_core.runnables.schema import StreamEvent
|
||||
from langchain_core.tools import BaseTool
|
||||
@@ -108,7 +108,7 @@ if TYPE_CHECKING:
|
||||
def _generate_response_from_error(error: BaseException) -> list[ChatGeneration]:
|
||||
if hasattr(error, "response"):
|
||||
response = error.response
|
||||
metadata: dict = {}
|
||||
metadata: dict[str, Any] = {}
|
||||
if hasattr(response, "json"):
|
||||
try:
|
||||
metadata["body"] = response.json()
|
||||
@@ -248,7 +248,9 @@ async def agenerate_from_stream(
|
||||
return await run_in_executor(None, generate_from_stream, iter(chunks))
|
||||
|
||||
|
||||
def _format_ls_structured_output(ls_structured_output_format: dict | None) -> dict:
|
||||
def _format_ls_structured_output(
|
||||
ls_structured_output_format: dict[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
if ls_structured_output_format:
|
||||
try:
|
||||
ls_structured_output_format_dict = {
|
||||
@@ -804,7 +806,7 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
|
||||
and isinstance(chunk.message, AIMessageChunk)
|
||||
and not chunk.message.chunk_position
|
||||
):
|
||||
empty_content: str | list = (
|
||||
empty_content: str | list[str | dict[str, Any]] = (
|
||||
"" if isinstance(chunk.message.content, str) else []
|
||||
)
|
||||
msg_chunk = AIMessageChunk(
|
||||
@@ -938,7 +940,7 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
|
||||
and isinstance(chunk.message, AIMessageChunk)
|
||||
and not chunk.message.chunk_position
|
||||
):
|
||||
empty_content: str | list = (
|
||||
empty_content: str | list[str | dict[str, Any]] = (
|
||||
"" if isinstance(chunk.message.content, str) else []
|
||||
)
|
||||
msg_chunk = AIMessageChunk(
|
||||
@@ -1372,11 +1374,13 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
|
||||
# --- Custom methods ---
|
||||
|
||||
def _combine_llm_outputs(
|
||||
self, _llm_outputs: list[builtins.dict | None], /
|
||||
) -> builtins.dict:
|
||||
self, _llm_outputs: list[builtins.dict[str, Any] | None], /
|
||||
) -> builtins.dict[str, Any]:
|
||||
return {}
|
||||
|
||||
def _convert_cached_generations(self, cache_val: list) -> list[ChatGeneration]:
|
||||
def _convert_cached_generations(
|
||||
self, cache_val: list[Generation]
|
||||
) -> list[ChatGeneration]:
|
||||
"""Convert cached Generation objects to ChatGeneration objects.
|
||||
|
||||
Handle case where cache contains Generation objects instead of
|
||||
@@ -1466,7 +1470,7 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
|
||||
self,
|
||||
stop: list[str] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> builtins.dict:
|
||||
) -> builtins.dict[str, Any]:
|
||||
params = self._dict_for_compat()
|
||||
params["stop"] = stop
|
||||
return {**params, **kwargs}
|
||||
@@ -1980,7 +1984,7 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
|
||||
and isinstance(chunk.message, AIMessageChunk)
|
||||
and not chunk.message.chunk_position
|
||||
):
|
||||
empty_content: str | list = (
|
||||
empty_content: str | list[str | dict[str, Any]] = (
|
||||
"" if isinstance(chunk.message.content, str) else []
|
||||
)
|
||||
chunk = ChatGenerationChunk(
|
||||
@@ -2137,7 +2141,7 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
|
||||
and isinstance(chunk.message, AIMessageChunk)
|
||||
and not chunk.message.chunk_position
|
||||
):
|
||||
empty_content: str | list = (
|
||||
empty_content: str | list[str | dict[str, Any]] = (
|
||||
"" if isinstance(chunk.message.content, str) else []
|
||||
)
|
||||
chunk = ChatGenerationChunk(
|
||||
@@ -2339,7 +2343,7 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
|
||||
|
||||
def bind_tools(
|
||||
self,
|
||||
tools: Sequence[builtins.dict[str, Any] | type | Callable | BaseTool],
|
||||
tools: Sequence[builtins.dict[str, Any] | type | Callable[..., Any] | BaseTool],
|
||||
*,
|
||||
tool_choice: str | None = None,
|
||||
**kwargs: Any,
|
||||
@@ -2519,8 +2523,9 @@ class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
|
||||
"schema": schema,
|
||||
},
|
||||
)
|
||||
output_parser: JsonOutputToolsParser
|
||||
if isinstance(schema, type) and is_basemodel_subclass(schema):
|
||||
output_parser: OutputParserLike = PydanticToolsParser(
|
||||
output_parser = PydanticToolsParser(
|
||||
tools=[cast("TypeBaseModel", schema)], first_tool_only=True
|
||||
)
|
||||
else:
|
||||
@@ -2679,7 +2684,7 @@ class SimpleChatModel(BaseChatModel):
|
||||
|
||||
def _gen_info_and_msg_metadata(
|
||||
generation: ChatGeneration | ChatGenerationChunk,
|
||||
) -> dict:
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
**(generation.generation_info or {}),
|
||||
**generation.message.response_metadata,
|
||||
|
||||
@@ -6,7 +6,7 @@ These are traditionally older models (newer models generally are chat models).
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import builtins # noqa: TC003
|
||||
import builtins
|
||||
import functools
|
||||
import inspect
|
||||
import json
|
||||
@@ -60,11 +60,12 @@ from langchain_core.runnables import RunnableConfig, ensure_config, get_config_l
|
||||
from langchain_core.runnables.config import run_in_executor
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import builtins
|
||||
import uuid
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_background_tasks: set[asyncio.Task] = set()
|
||||
_background_tasks: set[asyncio.Task[None]] = set()
|
||||
|
||||
|
||||
@functools.lru_cache
|
||||
@@ -159,7 +160,7 @@ def get_prompts(
|
||||
params: dict[str, Any],
|
||||
prompts: list[str],
|
||||
cache: BaseCache | bool | None = None, # noqa: FBT001
|
||||
) -> tuple[dict[int, list], str, list[int], list[str]]:
|
||||
) -> tuple[dict[int, list[Generation]], str, list[int], list[str]]:
|
||||
"""Get prompts that are already cached.
|
||||
|
||||
Args:
|
||||
@@ -195,7 +196,7 @@ async def aget_prompts(
|
||||
params: dict[str, Any],
|
||||
prompts: list[str],
|
||||
cache: BaseCache | bool | None = None, # noqa: FBT001
|
||||
) -> tuple[dict[int, list], str, list[int], list[str]]:
|
||||
) -> tuple[dict[int, list[Generation]], str, list[int], list[str]]:
|
||||
"""Get prompts that are already cached. Async version.
|
||||
|
||||
Args:
|
||||
@@ -228,12 +229,12 @@ async def aget_prompts(
|
||||
|
||||
def update_cache(
|
||||
cache: BaseCache | bool | None, # noqa: FBT001
|
||||
existing_prompts: dict[int, list],
|
||||
existing_prompts: dict[int, list[Generation]],
|
||||
llm_string: str,
|
||||
missing_prompt_idxs: list[int],
|
||||
new_results: LLMResult,
|
||||
prompts: list[str],
|
||||
) -> dict | None:
|
||||
) -> dict[str, Any] | None:
|
||||
"""Update the cache and get the LLM output.
|
||||
|
||||
Args:
|
||||
@@ -261,12 +262,12 @@ def update_cache(
|
||||
|
||||
async def aupdate_cache(
|
||||
cache: BaseCache | bool | None, # noqa: FBT001
|
||||
existing_prompts: dict[int, list],
|
||||
existing_prompts: dict[int, list[Generation]],
|
||||
llm_string: str,
|
||||
missing_prompt_idxs: list[int],
|
||||
new_results: LLMResult,
|
||||
prompts: list[str],
|
||||
) -> dict | None:
|
||||
) -> dict[str, Any] | None:
|
||||
"""Update the cache and get the LLM output. Async version.
|
||||
|
||||
Args:
|
||||
@@ -954,7 +955,8 @@ class BaseLLM(BaseLanguageModel[str], ABC):
|
||||
callbacks = cast("list[Callbacks]", callbacks)
|
||||
tags_list = cast("list[list[str] | None]", tags or ([None] * len(prompts)))
|
||||
metadata_list = cast(
|
||||
"list[dict[str, Any] | None]", metadata or ([{}] * len(prompts))
|
||||
"list[builtins.dict[str, Any] | None]",
|
||||
metadata or ([{}] * len(prompts)),
|
||||
)
|
||||
run_name_list = run_name or cast(
|
||||
"list[str | None]", ([None] * len(prompts))
|
||||
@@ -989,7 +991,7 @@ class BaseLLM(BaseLanguageModel[str], ABC):
|
||||
self.verbose,
|
||||
cast("list[str]", tags),
|
||||
self.tags,
|
||||
cast("dict[str, Any]", metadata),
|
||||
cast("builtins.dict[str, Any]", metadata),
|
||||
self.metadata,
|
||||
langsmith_inheritable_metadata=_filter_invocation_params_for_tracing(
|
||||
params
|
||||
@@ -1074,8 +1076,8 @@ class BaseLLM(BaseLanguageModel[str], ABC):
|
||||
|
||||
@staticmethod
|
||||
def _get_run_ids_list(
|
||||
run_id: uuid.UUID | list[uuid.UUID | None] | None, prompts: list
|
||||
) -> list:
|
||||
run_id: uuid.UUID | list[uuid.UUID | None] | None, prompts: list[str]
|
||||
) -> list[uuid.UUID | None]:
|
||||
if run_id is None:
|
||||
return [None] * len(prompts)
|
||||
if isinstance(run_id, list):
|
||||
@@ -1226,7 +1228,8 @@ class BaseLLM(BaseLanguageModel[str], ABC):
|
||||
callbacks = cast("list[Callbacks]", callbacks)
|
||||
tags_list = cast("list[list[str] | None]", tags or ([None] * len(prompts)))
|
||||
metadata_list = cast(
|
||||
"list[dict[str, Any] | None]", metadata or ([{}] * len(prompts))
|
||||
"list[builtins.dict[str, Any] | None]",
|
||||
metadata or ([{}] * len(prompts)),
|
||||
)
|
||||
run_name_list = run_name or cast(
|
||||
"list[str | None]", ([None] * len(prompts))
|
||||
@@ -1261,7 +1264,7 @@ class BaseLLM(BaseLanguageModel[str], ABC):
|
||||
self.verbose,
|
||||
cast("list[str]", tags),
|
||||
self.tags,
|
||||
cast("dict[str, Any]", metadata),
|
||||
cast("builtins.dict[str, Any]", metadata),
|
||||
self.metadata,
|
||||
langsmith_inheritable_metadata=_filter_invocation_params_for_tracing(
|
||||
params
|
||||
|
||||
@@ -165,7 +165,7 @@ class Serializable(BaseModel, ABC):
|
||||
return {}
|
||||
|
||||
@property
|
||||
def lc_attributes(self) -> dict:
|
||||
def lc_attributes(self) -> dict[str, Any]:
|
||||
"""List of attribute names that should be included in the serialized kwargs.
|
||||
|
||||
These attributes must be accepted by the constructor.
|
||||
|
||||
@@ -185,21 +185,21 @@ class AIMessage(BaseMessage):
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict],
|
||||
content: str | list[str | dict[Any, Any]],
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict] | None = None,
|
||||
content: str | list[str | dict[Any, Any]] | None = None,
|
||||
content_blocks: list[types.ContentBlock] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict] | None = None,
|
||||
content: str | list[str | dict[Any, Any]] | None = None,
|
||||
content_blocks: list[types.ContentBlock] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
@@ -221,14 +221,14 @@ class AIMessage(BaseMessage):
|
||||
kwargs["tool_calls"] = content_tool_calls
|
||||
|
||||
super().__init__(
|
||||
content=cast("str | list[str | dict]", content_blocks),
|
||||
content=cast("list[str | dict[Any, Any]]", content_blocks),
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
super().__init__(content=content, **kwargs)
|
||||
|
||||
@property
|
||||
def lc_attributes(self) -> dict:
|
||||
def lc_attributes(self) -> dict[str, Any]:
|
||||
"""Attributes to be serialized.
|
||||
|
||||
Includes all attributes, even if they are derived from other initialization
|
||||
@@ -301,7 +301,7 @@ class AIMessage(BaseMessage):
|
||||
# TODO: remove this logic if possible, reducing breaking nature of changes
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def _backwards_compat_tool_calls(cls, values: dict) -> Any:
|
||||
def _backwards_compat_tool_calls(cls, values: dict[str, Any]) -> Any:
|
||||
check_additional_kwargs = not any(
|
||||
values.get(k)
|
||||
for k in ("tool_calls", "invalid_tool_calls", "tool_call_chunks")
|
||||
@@ -431,7 +431,7 @@ class AIMessageChunk(AIMessage, BaseMessageChunk):
|
||||
|
||||
@property
|
||||
@override
|
||||
def lc_attributes(self) -> dict:
|
||||
def lc_attributes(self) -> dict[str, Any]:
|
||||
return {
|
||||
"tool_calls": self.tool_calls,
|
||||
"invalid_tool_calls": self.invalid_tool_calls,
|
||||
@@ -769,8 +769,8 @@ def add_usage(left: UsageMetadata | None, right: UsageMetadata | None) -> UsageM
|
||||
**cast(
|
||||
"UsageMetadata",
|
||||
_dict_int_op(
|
||||
cast("dict", left),
|
||||
cast("dict", right),
|
||||
cast("dict[str, Any]", left),
|
||||
cast("dict[str, Any]", right),
|
||||
operator.add,
|
||||
),
|
||||
)
|
||||
@@ -832,8 +832,8 @@ def subtract_usage(
|
||||
**cast(
|
||||
"UsageMetadata",
|
||||
_dict_int_op(
|
||||
cast("dict", left),
|
||||
cast("dict", right),
|
||||
cast("dict[str, Any]", left),
|
||||
cast("dict[str, Any]", right),
|
||||
(lambda le, ri: max(le - ri, 0)),
|
||||
),
|
||||
)
|
||||
|
||||
@@ -100,10 +100,10 @@ class BaseMessage(Serializable):
|
||||
[`SystemMessage`][langchain.messages.SystemMessage].
|
||||
"""
|
||||
|
||||
content: str | list[str | dict]
|
||||
content: str | list[str | dict[Any, Any]]
|
||||
"""The contents of the message."""
|
||||
|
||||
additional_kwargs: dict = Field(default_factory=dict)
|
||||
additional_kwargs: dict[Any, Any] = Field(default_factory=dict)
|
||||
"""Reserved for additional payload data associated with the message.
|
||||
|
||||
For example, for a message from an AI, this could include tool calls as
|
||||
@@ -111,7 +111,7 @@ class BaseMessage(Serializable):
|
||||
|
||||
"""
|
||||
|
||||
response_metadata: dict = Field(default_factory=dict)
|
||||
response_metadata: dict[Any, Any] = Field(default_factory=dict)
|
||||
"""Examples: response headers, logprobs, token counts, model name."""
|
||||
|
||||
type: str
|
||||
@@ -146,21 +146,21 @@ class BaseMessage(Serializable):
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict],
|
||||
content: str | list[str | dict[Any, Any]],
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict] | None = None,
|
||||
content: str | list[str | dict[Any, Any]] | None = None,
|
||||
content_blocks: list[types.ContentBlock] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict] | None = None,
|
||||
content: str | list[str | dict[Any, Any]] | None = None,
|
||||
content_blocks: list[types.ContentBlock] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
@@ -364,9 +364,9 @@ class BaseMessage(Serializable):
|
||||
|
||||
|
||||
def merge_content(
|
||||
first_content: str | list[str | dict],
|
||||
*contents: str | list[str | dict],
|
||||
) -> str | list[str | dict]:
|
||||
first_content: str | list[str | dict[Any, Any]],
|
||||
*contents: str | list[str | dict[Any, Any]],
|
||||
) -> str | list[str | dict[Any, Any]]:
|
||||
"""Merge multiple message contents.
|
||||
|
||||
Args:
|
||||
@@ -377,7 +377,7 @@ def merge_content(
|
||||
The merged content.
|
||||
|
||||
"""
|
||||
merged: str | list[str | dict]
|
||||
merged: str | list[str | dict[Any, Any]]
|
||||
merged = "" if first_content is None else first_content
|
||||
|
||||
for content in contents:
|
||||
@@ -391,7 +391,7 @@ def merge_content(
|
||||
merged = [merged, *content]
|
||||
elif isinstance(content, list):
|
||||
# If both are lists
|
||||
merged = merge_lists(cast("list", merged), content) # type: ignore[assignment]
|
||||
merged = merge_lists(merged, content) # type: ignore[assignment]
|
||||
# If the first content is a list, and the second content is a string
|
||||
# If the last element of the first content is a string
|
||||
# Add the second content to the last element
|
||||
@@ -471,7 +471,7 @@ class BaseMessageChunk(BaseMessage):
|
||||
raise TypeError(msg)
|
||||
|
||||
|
||||
def message_to_dict(message: BaseMessage) -> dict:
|
||||
def message_to_dict(message: BaseMessage) -> dict[str, Any]:
|
||||
"""Convert a Message to a dictionary.
|
||||
|
||||
Args:
|
||||
@@ -485,7 +485,7 @@ def message_to_dict(message: BaseMessage) -> dict:
|
||||
return {"type": message.type, "data": message.model_dump()}
|
||||
|
||||
|
||||
def messages_to_dict(messages: Sequence[BaseMessage]) -> list[dict]:
|
||||
def messages_to_dict(messages: Sequence[BaseMessage]) -> list[dict[str, Any]]:
|
||||
"""Convert a sequence of Messages to a list of dictionaries.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -197,8 +197,9 @@ def _convert_citation_to_v1(citation: dict[str, Any]) -> types.Annotation:
|
||||
|
||||
def _convert_to_v1_from_anthropic(message: AIMessage) -> list[types.ContentBlock]:
|
||||
"""Convert Anthropic message content to v1 format."""
|
||||
content: list[str | dict[str, Any]]
|
||||
if isinstance(message.content, str):
|
||||
content: list[str | dict] = [{"type": "text", "text": message.content}]
|
||||
content = [{"type": "text", "text": message.content}]
|
||||
else:
|
||||
content = message.content
|
||||
|
||||
|
||||
@@ -27,7 +27,7 @@ def _populate_extras(
|
||||
return standard_block
|
||||
|
||||
|
||||
def _parse_code_json(s: str) -> dict:
|
||||
def _parse_code_json(s: str) -> dict[str, Any]:
|
||||
"""Extract Python code from Groq built-in tool content.
|
||||
|
||||
Extracts the value of the 'code' field from a string of the form:
|
||||
|
||||
@@ -44,8 +44,8 @@ def _convert_v0_multimodal_input_to_v1(
|
||||
|
||||
|
||||
def _convert_legacy_v0_content_block_to_v1(
|
||||
block: dict,
|
||||
) -> types.ContentBlock | dict:
|
||||
block: dict[str, Any],
|
||||
) -> types.ContentBlock | dict[str, Any]:
|
||||
"""Convert a LangChain v0 content block to v1 format.
|
||||
|
||||
Preserves unknown keys as extras to avoid data loss.
|
||||
@@ -53,7 +53,9 @@ def _convert_legacy_v0_content_block_to_v1(
|
||||
Returns the original block unchanged if it's not in v0 format.
|
||||
"""
|
||||
|
||||
def _extract_v0_extras(block_dict: dict, known_keys: set[str]) -> dict[str, Any]:
|
||||
def _extract_v0_extras(
|
||||
block_dict: dict[str, Any], known_keys: set[str]
|
||||
) -> dict[str, Any]:
|
||||
"""Extract unknown keys from v0 block to preserve as extras.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -19,7 +19,7 @@ if TYPE_CHECKING:
|
||||
from langchain_core.messages import AIMessage
|
||||
|
||||
|
||||
def convert_to_openai_image_block(block: dict[str, Any]) -> dict:
|
||||
def convert_to_openai_image_block(block: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Convert `ImageContentBlock` to format expected by OpenAI Chat Completions.
|
||||
|
||||
Args:
|
||||
@@ -56,8 +56,9 @@ def convert_to_openai_image_block(block: dict[str, Any]) -> dict:
|
||||
|
||||
|
||||
def convert_to_openai_data_block(
|
||||
block: dict, api: Literal["chat/completions", "responses"] = "chat/completions"
|
||||
) -> dict:
|
||||
block: dict[str, Any],
|
||||
api: Literal["chat/completions", "responses"] = "chat/completions",
|
||||
) -> dict[str, Any]:
|
||||
"""Format standard data content block to format expected by OpenAI.
|
||||
|
||||
"Standard data content block" can include old-style LangChain v0 blocks
|
||||
@@ -265,7 +266,7 @@ def _convert_to_v1_from_chat_completions_chunk(
|
||||
def _convert_from_v1_to_chat_completions(message: AIMessage) -> AIMessage:
|
||||
"""Convert a v1 message to the Chat Completions format."""
|
||||
if isinstance(message.content, list):
|
||||
new_content: list = []
|
||||
new_content: list[Any] = []
|
||||
for block in message.content:
|
||||
if isinstance(block, dict):
|
||||
block_type = block.get("type")
|
||||
@@ -330,7 +331,7 @@ def _convert_from_v03_ai_message(message: AIMessage) -> AIMessage:
|
||||
]
|
||||
|
||||
# Build a bucket for every known block type
|
||||
buckets: dict[str, list] = {key: [] for key in content_order}
|
||||
buckets: dict[str, list[Any]] = {key: [] for key in content_order}
|
||||
unknown_blocks = []
|
||||
|
||||
# Reasoning
|
||||
@@ -422,8 +423,8 @@ def _convert_from_v03_ai_message(message: AIMessage) -> AIMessage:
|
||||
|
||||
|
||||
def _convert_openai_format_to_data_block(
|
||||
block: dict,
|
||||
) -> types.ContentBlock | dict[Any, Any]:
|
||||
block: dict[str, Any],
|
||||
) -> types.ContentBlock | dict[str, Any]:
|
||||
"""Convert OpenAI image/audio/file content block to respective v1 multimodal block.
|
||||
|
||||
We expect that the incoming block is verified to be in OpenAI Chat Completions
|
||||
@@ -439,7 +440,9 @@ def _convert_openai_format_to_data_block(
|
||||
"""
|
||||
|
||||
# Extract extra keys to put them in `extras`
|
||||
def _extract_extras(block_dict: dict, known_keys: set[str]) -> dict[str, Any]:
|
||||
def _extract_extras(
|
||||
block_dict: dict[str, Any], known_keys: set[str]
|
||||
) -> dict[str, Any]:
|
||||
"""Extract unknown keys from block to preserve as extras."""
|
||||
return {k: v for k, v in block_dict.items() if k not in known_keys}
|
||||
|
||||
|
||||
@@ -905,7 +905,7 @@ def _get_data_content_block_types() -> tuple[str, ...]:
|
||||
return tuple(data_block_types)
|
||||
|
||||
|
||||
def is_data_content_block(block: dict) -> bool:
|
||||
def is_data_content_block(block: dict[str, Any]) -> bool:
|
||||
"""Check if the provided content block is a data content block.
|
||||
|
||||
Returns True for both v0 (old-style) and v1 (new-style) multimodal data blocks.
|
||||
|
||||
@@ -32,28 +32,28 @@ class HumanMessage(BaseMessage):
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict],
|
||||
content: str | list[str | dict[Any, Any]],
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict] | None = None,
|
||||
content: str | list[str | dict[Any, Any]] | None = None,
|
||||
content_blocks: list[types.ContentBlock] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict] | None = None,
|
||||
content: str | list[str | dict[Any, Any]] | None = None,
|
||||
content_blocks: list[types.ContentBlock] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Specify `content` as positional arg or `content_blocks` for typing."""
|
||||
if content_blocks is not None:
|
||||
super().__init__(
|
||||
content=cast("str | list[str | dict]", content_blocks),
|
||||
content=cast("list[str | dict[Any, Any]]", content_blocks),
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
|
||||
@@ -32,28 +32,28 @@ class SystemMessage(BaseMessage):
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict],
|
||||
content: str | list[str | dict[Any, Any]],
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict] | None = None,
|
||||
content: str | list[str | dict[Any, Any]] | None = None,
|
||||
content_blocks: list[types.ContentBlock] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict] | None = None,
|
||||
content: str | list[str | dict[Any, Any]] | None = None,
|
||||
content_blocks: list[types.ContentBlock] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Specify `content` as positional arg or `content_blocks` for typing."""
|
||||
if content_blocks is not None:
|
||||
super().__init__(
|
||||
content=cast("str | list[str | dict]", content_blocks),
|
||||
content=cast("list[str | dict[Any, Any]]", content_blocks),
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
|
||||
@@ -76,20 +76,20 @@ class ToolMessage(BaseMessage, ToolOutputMixin):
|
||||
Should only be specified if it is different from the message content, e.g. if only
|
||||
a subset of the full tool output is being passed as message content but the full
|
||||
output is needed in other parts of the code.
|
||||
|
||||
"""
|
||||
|
||||
status: Literal["success", "error"] = "success"
|
||||
"""Status of the tool invocation."""
|
||||
|
||||
additional_kwargs: dict = Field(default_factory=dict, repr=False)
|
||||
additional_kwargs: dict[Any, Any] = Field(default_factory=dict, repr=False)
|
||||
"""Currently inherited from `BaseMessage`, but not used."""
|
||||
response_metadata: dict = Field(default_factory=dict, repr=False)
|
||||
|
||||
response_metadata: dict[Any, Any] = Field(default_factory=dict, repr=False)
|
||||
"""Currently inherited from `BaseMessage`, but not used."""
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def coerce_args(cls, values: dict) -> dict:
|
||||
def coerce_args(cls, values: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Coerce the model arguments to the correct types.
|
||||
|
||||
Args:
|
||||
@@ -135,21 +135,21 @@ class ToolMessage(BaseMessage, ToolOutputMixin):
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict],
|
||||
content: str | list[str | dict[Any, Any]],
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict] | None = None,
|
||||
content: str | list[str | dict[Any, Any]] | None = None,
|
||||
content_blocks: list[types.ContentBlock] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
content: str | list[str | dict] | None = None,
|
||||
content: str | list[str | dict[Any, Any]] | None = None,
|
||||
content_blocks: list[types.ContentBlock] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
@@ -164,7 +164,7 @@ class ToolMessage(BaseMessage, ToolOutputMixin):
|
||||
"""
|
||||
if content_blocks is not None:
|
||||
super().__init__(
|
||||
content=cast("str | list[str | dict]", content_blocks),
|
||||
content=cast("list[str | dict[Any, Any]]", content_blocks),
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
@@ -347,7 +347,7 @@ def invalid_tool_call(
|
||||
|
||||
|
||||
def default_tool_parser(
|
||||
raw_tool_calls: list[dict],
|
||||
raw_tool_calls: list[dict[str, Any]],
|
||||
) -> tuple[list[ToolCall], list[InvalidToolCall]]:
|
||||
"""Best-effort parsing of tools.
|
||||
|
||||
@@ -383,7 +383,9 @@ def default_tool_parser(
|
||||
return tool_calls, invalid_tool_calls
|
||||
|
||||
|
||||
def default_tool_chunk_parser(raw_tool_calls: list[dict]) -> list[ToolCallChunk]:
|
||||
def default_tool_chunk_parser(
|
||||
raw_tool_calls: list[dict[str, Any]],
|
||||
) -> list[ToolCallChunk]:
|
||||
"""Best-effort parsing of tool chunks.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -512,7 +512,7 @@ def get_buffer_string(
|
||||
return message_separator.join(string_messages)
|
||||
|
||||
|
||||
def _message_from_dict(message: dict) -> BaseMessage:
|
||||
def _message_from_dict(message: dict[str, Any]) -> BaseMessage:
|
||||
type_ = message["type"]
|
||||
if type_ == "human":
|
||||
return HumanMessage(**message["data"])
|
||||
@@ -544,7 +544,7 @@ def _message_from_dict(message: dict) -> BaseMessage:
|
||||
raise ValueError(msg)
|
||||
|
||||
|
||||
def messages_from_dict(messages: Sequence[dict]) -> list[BaseMessage]:
|
||||
def messages_from_dict(messages: Sequence[dict[str, Any]]) -> list[BaseMessage]:
|
||||
"""Convert a sequence of messages from dicts to `Message` objects.
|
||||
|
||||
Args:
|
||||
@@ -709,9 +709,9 @@ def _convert_to_message(message: MessageLikeRepresentation) -> BaseMessage:
|
||||
- 2-tuple of (role string, template); e.g., (`'human'`, `'{user_input}'`)
|
||||
- dict: a message dict with role and content keys
|
||||
- dict: the `Serializable` constructor-envelope wire shape
|
||||
`{"lc": 1, "type": "constructor", "id": [..., "<ClassName>"],
|
||||
"kwargs": {...}}` — unpacked structurally and routed through the
|
||||
standard dict-with-type dispatch.
|
||||
`{"lc": 1, "type": "constructor", "id": [..., "<ClassName>"],
|
||||
"kwargs": {...}}` — unpacked structurally and routed through the
|
||||
standard dict-with-type dispatch.
|
||||
- string: shorthand for (`'human'`, template); e.g., `'{user_input}'`
|
||||
|
||||
Args:
|
||||
@@ -1132,7 +1132,7 @@ def trim_messages(
|
||||
max_tokens: int,
|
||||
token_counter: Callable[[list[BaseMessage]], int]
|
||||
| Callable[[BaseMessage], int]
|
||||
| BaseLanguageModel
|
||||
| BaseLanguageModel[Any]
|
||||
| Literal["approximate"],
|
||||
strategy: Literal["first", "last"] = "last",
|
||||
allow_partial: bool = False,
|
||||
@@ -1484,10 +1484,11 @@ def trim_messages(
|
||||
)
|
||||
raise ValueError(msg)
|
||||
|
||||
text_splitter_fn: Callable[[str], list[str]]
|
||||
if _HAS_LANGCHAIN_TEXT_SPLITTERS and isinstance(text_splitter, TextSplitter):
|
||||
text_splitter_fn = text_splitter.split_text
|
||||
elif text_splitter:
|
||||
text_splitter_fn = cast("Callable", text_splitter)
|
||||
text_splitter_fn = cast("Callable[[str], list[str]]", text_splitter)
|
||||
else:
|
||||
text_splitter_fn = _default_text_splitter
|
||||
|
||||
@@ -1528,7 +1529,7 @@ def convert_to_openai_messages(
|
||||
text_format: Literal["string", "block"] = "string",
|
||||
include_id: bool = False,
|
||||
pass_through_unknown_blocks: bool = True,
|
||||
) -> dict: ...
|
||||
) -> dict[str, Any]: ...
|
||||
|
||||
|
||||
@overload
|
||||
@@ -1538,7 +1539,7 @@ def convert_to_openai_messages(
|
||||
text_format: Literal["string", "block"] = "string",
|
||||
include_id: bool = False,
|
||||
pass_through_unknown_blocks: bool = True,
|
||||
) -> list[dict]: ...
|
||||
) -> list[dict[str, Any]]: ...
|
||||
|
||||
|
||||
def convert_to_openai_messages(
|
||||
@@ -1547,7 +1548,7 @@ def convert_to_openai_messages(
|
||||
text_format: Literal["string", "block"] = "string",
|
||||
include_id: bool = False,
|
||||
pass_through_unknown_blocks: bool = True,
|
||||
) -> dict | list[dict]:
|
||||
) -> dict[str, Any] | list[dict[str, Any]]:
|
||||
"""Convert LangChain messages into OpenAI message dicts.
|
||||
|
||||
Args:
|
||||
@@ -1636,7 +1637,7 @@ def convert_to_openai_messages(
|
||||
err = f"Unrecognized {text_format=}, expected one of 'string' or 'block'."
|
||||
raise ValueError(err)
|
||||
|
||||
oai_messages: list[dict] = []
|
||||
oai_messages: list[dict[str, Any]] = []
|
||||
|
||||
if is_single := isinstance(messages, (BaseMessage, dict, str)):
|
||||
messages = [messages]
|
||||
@@ -1644,9 +1645,9 @@ def convert_to_openai_messages(
|
||||
messages = convert_to_messages(messages)
|
||||
|
||||
for i, message in enumerate(messages):
|
||||
oai_msg: dict = {"role": _get_message_openai_role(message)}
|
||||
tool_messages: list = []
|
||||
content: str | list[dict]
|
||||
oai_msg: dict[str, Any] = {"role": _get_message_openai_role(message)}
|
||||
tool_messages: list[dict[str, Any]] = []
|
||||
content: str | list[dict[str, Any]]
|
||||
|
||||
if message.name:
|
||||
oai_msg["name"] = message.name
|
||||
@@ -2216,7 +2217,7 @@ def _get_message_openai_role(message: BaseMessage) -> str:
|
||||
raise ValueError(msg)
|
||||
|
||||
|
||||
def _convert_to_openai_tool_calls(tool_calls: list[ToolCall]) -> list[dict]:
|
||||
def _convert_to_openai_tool_calls(tool_calls: list[ToolCall]) -> list[dict[str, Any]]:
|
||||
return [
|
||||
{
|
||||
"type": "function",
|
||||
@@ -2247,7 +2248,7 @@ def count_tokens_approximately(
|
||||
- For AI messages, the token count also includes stringified tool calls.
|
||||
- For tool messages, the token count also includes the tool call ID.
|
||||
- For multimodal messages with images, applies a fixed token penalty per image
|
||||
instead of counting base64-encoded characters.
|
||||
instead of counting base64-encoded characters.
|
||||
- If tools are provided, the token count also includes stringified tool schemas.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import builtins # noqa: TC003
|
||||
import builtins
|
||||
import contextlib
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import (
|
||||
@@ -23,6 +23,8 @@ from langchain_core.runnables import Runnable, RunnableConfig, RunnableSerializa
|
||||
from langchain_core.runnables.config import run_in_executor
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import builtins
|
||||
|
||||
from langchain_core.prompt_values import PromptValue
|
||||
|
||||
T = TypeVar("T")
|
||||
@@ -70,7 +72,7 @@ class BaseLLMOutputParser(ABC, Generic[T]):
|
||||
|
||||
|
||||
class BaseGenerationOutputParser(
|
||||
BaseLLMOutputParser, RunnableSerializable[LanguageModelOutput, T]
|
||||
BaseLLMOutputParser[T], RunnableSerializable[LanguageModelOutput, T]
|
||||
):
|
||||
"""Base class to parse the output of an LLM call."""
|
||||
|
||||
@@ -136,7 +138,7 @@ class BaseGenerationOutputParser(
|
||||
|
||||
|
||||
class BaseOutputParser(
|
||||
BaseLLMOutputParser, RunnableSerializable[LanguageModelOutput, T]
|
||||
BaseLLMOutputParser[T], RunnableSerializable[LanguageModelOutput, T]
|
||||
):
|
||||
"""Base class to parse the output of an LLM call.
|
||||
|
||||
|
||||
@@ -58,7 +58,7 @@ class ListOutputParser(BaseTransformOutputParser[list[str]]):
|
||||
A list of strings.
|
||||
"""
|
||||
|
||||
def parse_iter(self, text: str) -> Iterator[re.Match]:
|
||||
def parse_iter(self, text: str) -> Iterator[re.Match[str]]:
|
||||
"""Parse the output of an LLM call.
|
||||
|
||||
Args:
|
||||
@@ -210,7 +210,7 @@ class NumberedListOutputParser(ListOutputParser):
|
||||
return re.findall(self.pattern, text)
|
||||
|
||||
@override
|
||||
def parse_iter(self, text: str) -> Iterator[re.Match]:
|
||||
def parse_iter(self, text: str) -> Iterator[re.Match[str]]:
|
||||
return re.finditer(self.pattern, text)
|
||||
|
||||
@property
|
||||
@@ -241,7 +241,7 @@ class MarkdownListOutputParser(ListOutputParser):
|
||||
return re.findall(self.pattern, text, re.MULTILINE)
|
||||
|
||||
@override
|
||||
def parse_iter(self, text: str) -> Iterator[re.Match]:
|
||||
def parse_iter(self, text: str) -> Iterator[re.Match[str]]:
|
||||
return re.finditer(self.pattern, text, re.MULTILINE)
|
||||
|
||||
@property
|
||||
|
||||
@@ -100,7 +100,7 @@ def make_invalid_tool_call(
|
||||
|
||||
|
||||
def parse_tool_calls(
|
||||
raw_tool_calls: list[dict],
|
||||
raw_tool_calls: list[dict[str, Any]],
|
||||
*,
|
||||
partial: bool = False,
|
||||
strict: bool = False,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""Output parsers using Pydantic."""
|
||||
|
||||
import json
|
||||
from typing import Annotated, Generic, Literal, overload
|
||||
from typing import Annotated, Any, Generic, Literal, overload
|
||||
|
||||
import pydantic
|
||||
from pydantic import SkipValidation
|
||||
@@ -22,7 +22,7 @@ class PydanticOutputParser(JsonOutputParser, Generic[TBaseModel]):
|
||||
pydantic_object: Annotated[type[TBaseModel], SkipValidation()]
|
||||
"""The Pydantic model to parse."""
|
||||
|
||||
def _parse_obj(self, obj: dict) -> TBaseModel:
|
||||
def _parse_obj(self, obj: Any) -> TBaseModel:
|
||||
try:
|
||||
if issubclass(self.pydantic_object, pydantic.BaseModel):
|
||||
return self.pydantic_object.model_validate(obj)
|
||||
@@ -35,7 +35,7 @@ class PydanticOutputParser(JsonOutputParser, Generic[TBaseModel]):
|
||||
raise self._parser_exception(e, obj) from e
|
||||
|
||||
def _parser_exception(
|
||||
self, e: Exception, json_object: dict
|
||||
self, e: Exception, json_object: Any
|
||||
) -> OutputParserException:
|
||||
json_string = json.dumps(json_object, ensure_ascii=False)
|
||||
name = self.pydantic_object.__name__
|
||||
|
||||
@@ -148,7 +148,7 @@ class _StreamingParser:
|
||||
self.pull_parser.close()
|
||||
|
||||
|
||||
class XMLOutputParser(BaseTransformOutputParser):
|
||||
class XMLOutputParser(BaseTransformOutputParser[dict[str, Any]]):
|
||||
"""Parse an output using xml format.
|
||||
|
||||
Returns a dictionary of tags.
|
||||
@@ -170,7 +170,7 @@ class XMLOutputParser(BaseTransformOutputParser):
|
||||
3. A badly-formatted XML instance (unexpected 'tag' element):
|
||||
`'<foo>\n <tag>\n </tag>\n</foo>'`
|
||||
"""
|
||||
encoding_matcher: re.Pattern = re.compile(
|
||||
encoding_matcher: re.Pattern[str] = re.compile(
|
||||
r"<([^>]*encoding[^>]*)>\n(.*)", re.MULTILINE | re.DOTALL
|
||||
)
|
||||
|
||||
@@ -272,13 +272,13 @@ class XMLOutputParser(BaseTransformOutputParser):
|
||||
# If root text contains any non-whitespace character it
|
||||
# returns {root.tag: root.text}
|
||||
return {root.tag: root.text}
|
||||
result: dict = {root.tag: []}
|
||||
root_tag: list[Any] = []
|
||||
for child in root:
|
||||
if len(child) == 0:
|
||||
result[root.tag].append({child.tag: child.text})
|
||||
root_tag.append({child.tag: child.text})
|
||||
else:
|
||||
result[root.tag].append(self._root_to_dict(child))
|
||||
return result
|
||||
root_tag.append(self._root_to_dict(child))
|
||||
return {root.tag: root_tag}
|
||||
|
||||
@property
|
||||
def _type(self) -> str:
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""Chat result schema."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from langchain_core.outputs.chat_generation import ChatGeneration
|
||||
@@ -25,7 +27,7 @@ class ChatResult(BaseModel):
|
||||
input prompt.
|
||||
"""
|
||||
|
||||
llm_output: dict | None = None
|
||||
llm_output: dict[str, Any] | None = None
|
||||
"""For arbitrary model provider-specific output.
|
||||
|
||||
This dictionary is a free-form dictionary that can contain any information that the
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from copy import deepcopy
|
||||
from typing import Literal
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
@@ -37,7 +37,7 @@ class LLMResult(BaseModel):
|
||||
chat message.
|
||||
"""
|
||||
|
||||
llm_output: dict | None = None
|
||||
llm_output: dict[str, Any] | None = None
|
||||
"""For arbitrary model provider-specific output.
|
||||
|
||||
This dictionary is a free-form dictionary that can contain any information that the
|
||||
|
||||
@@ -8,7 +8,7 @@ from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Sequence
|
||||
from typing import Literal, cast
|
||||
from typing import Any, Literal, cast
|
||||
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
@@ -146,7 +146,7 @@ class ImagePromptValue(PromptValue):
|
||||
|
||||
def to_messages(self) -> list[BaseMessage]:
|
||||
"""Return prompt (image URL) as messages."""
|
||||
return [HumanMessage(content=[cast("dict", self.image_url)])]
|
||||
return [HumanMessage(content=[cast("dict[str, Any]", self.image_url)])]
|
||||
|
||||
|
||||
class ChatPromptValueConcrete(ChatPromptValue):
|
||||
|
||||
@@ -36,7 +36,7 @@ FormatOutputType = TypeVar("FormatOutputType")
|
||||
|
||||
|
||||
class BasePromptTemplate(
|
||||
RunnableSerializable[dict, PromptValue], ABC, Generic[FormatOutputType]
|
||||
RunnableSerializable[dict[str, Any], PromptValue], ABC, Generic[FormatOutputType]
|
||||
):
|
||||
"""Base class for all prompt templates, returning a prompt."""
|
||||
|
||||
@@ -155,7 +155,7 @@ class BasePromptTemplate(
|
||||
field_definitions={**required_input_variables, **optional_input_variables},
|
||||
)
|
||||
|
||||
def _validate_input(self, inner_input: Any) -> builtins.dict:
|
||||
def _validate_input(self, inner_input: Any) -> builtins.dict[str, Any]:
|
||||
if not isinstance(inner_input, dict):
|
||||
if len(self.input_variables) == 1:
|
||||
var_name = self.input_variables[0]
|
||||
@@ -192,22 +192,23 @@ class BasePromptTemplate(
|
||||
return inner_input_
|
||||
|
||||
def _format_prompt_with_error_handling(
|
||||
self,
|
||||
inner_input: builtins.dict,
|
||||
self, inner_input: builtins.dict[str, Any]
|
||||
) -> PromptValue:
|
||||
inner_input_ = self._validate_input(inner_input)
|
||||
return self.format_prompt(**inner_input_)
|
||||
|
||||
async def _aformat_prompt_with_error_handling(
|
||||
self,
|
||||
inner_input: builtins.dict,
|
||||
self, inner_input: builtins.dict[str, Any]
|
||||
) -> PromptValue:
|
||||
inner_input_ = self._validate_input(inner_input)
|
||||
return await self.aformat_prompt(**inner_input_)
|
||||
|
||||
@override
|
||||
def invoke(
|
||||
self, input: builtins.dict, config: RunnableConfig | None = None, **kwargs: Any
|
||||
self,
|
||||
input: builtins.dict[str, Any],
|
||||
config: RunnableConfig | None = None,
|
||||
**kwargs: Any,
|
||||
) -> PromptValue:
|
||||
"""Invoke the prompt.
|
||||
|
||||
@@ -233,7 +234,10 @@ class BasePromptTemplate(
|
||||
|
||||
@override
|
||||
async def ainvoke(
|
||||
self, input: builtins.dict, config: RunnableConfig | None = None, **kwargs: Any
|
||||
self,
|
||||
input: builtins.dict[str, Any],
|
||||
config: RunnableConfig | None = None,
|
||||
**kwargs: Any,
|
||||
) -> PromptValue:
|
||||
"""Async invoke the prompt.
|
||||
|
||||
@@ -279,7 +283,9 @@ class BasePromptTemplate(
|
||||
"""
|
||||
return self.format_prompt(**kwargs)
|
||||
|
||||
def partial(self, **kwargs: str | Callable[[], str]) -> BasePromptTemplate:
|
||||
def partial(
|
||||
self, **kwargs: str | Callable[[], str]
|
||||
) -> BasePromptTemplate[FormatOutputType]:
|
||||
"""Return a partial of the prompt template.
|
||||
|
||||
Args:
|
||||
@@ -420,7 +426,9 @@ class BasePromptTemplate(
|
||||
raise ValueError(msg)
|
||||
|
||||
|
||||
def _get_document_info(doc: Document, prompt: BasePromptTemplate[str]) -> dict:
|
||||
def _get_document_info(
|
||||
doc: Document, prompt: BasePromptTemplate[str]
|
||||
) -> dict[str, Any]:
|
||||
base_info = {"page_content": doc.page_content, **doc.metadata}
|
||||
missing_metadata = set(prompt.input_variables).difference(base_info)
|
||||
if len(missing_metadata) > 0:
|
||||
|
||||
@@ -33,7 +33,7 @@ from langchain_core.messages import (
|
||||
convert_to_messages,
|
||||
)
|
||||
from langchain_core.messages.base import get_msg_title_repr
|
||||
from langchain_core.prompt_values import ChatPromptValue, ImageURL
|
||||
from langchain_core.prompt_values import ChatPromptValue
|
||||
from langchain_core.prompts.base import BasePromptTemplate
|
||||
from langchain_core.prompts.dict import DictPromptTemplate
|
||||
from langchain_core.prompts.image import ImagePromptTemplate
|
||||
@@ -229,7 +229,7 @@ class BaseStringMessagePromptTemplate(BaseMessagePromptTemplate, ABC):
|
||||
prompt: StringPromptTemplate
|
||||
"""String prompt template."""
|
||||
|
||||
additional_kwargs: dict = Field(default_factory=dict)
|
||||
additional_kwargs: dict[str, Any] = Field(default_factory=dict)
|
||||
"""Additional keyword arguments to pass to the prompt template."""
|
||||
|
||||
@classmethod
|
||||
@@ -387,11 +387,11 @@ class ChatMessagePromptTemplate(BaseStringMessagePromptTemplate):
|
||||
|
||||
|
||||
class _TextTemplateParam(TypedDict, total=False):
|
||||
text: str | dict
|
||||
text: str | dict[str, Any]
|
||||
|
||||
|
||||
class _ImageTemplateParam(TypedDict, total=False):
|
||||
image_url: str | dict
|
||||
image_url: str | dict[str, Any]
|
||||
|
||||
|
||||
class _StringImageMessagePromptTemplate(BaseMessagePromptTemplate):
|
||||
@@ -402,7 +402,7 @@ class _StringImageMessagePromptTemplate(BaseMessagePromptTemplate):
|
||||
| list[StringPromptTemplate | ImagePromptTemplate | DictPromptTemplate]
|
||||
)
|
||||
"""Prompt template."""
|
||||
additional_kwargs: dict = Field(default_factory=dict)
|
||||
additional_kwargs: dict[str, Any] = Field(default_factory=dict)
|
||||
"""Additional keyword arguments to pass to the prompt template."""
|
||||
|
||||
_msg_class: type[BaseMessage]
|
||||
@@ -411,7 +411,7 @@ class _StringImageMessagePromptTemplate(BaseMessagePromptTemplate):
|
||||
def from_template(
|
||||
cls: type[Self],
|
||||
template: str
|
||||
| list[str | _TextTemplateParam | _ImageTemplateParam | dict[str, Any]],
|
||||
| Sequence[str | _TextTemplateParam | _ImageTemplateParam | dict[str, Any]],
|
||||
template_format: PromptTemplateFormat = "f-string",
|
||||
*,
|
||||
partial_variables: dict[str, Any] | None = None,
|
||||
@@ -434,14 +434,18 @@ class _StringImageMessagePromptTemplate(BaseMessagePromptTemplate):
|
||||
Raises:
|
||||
ValueError: If the template is not a string or list of strings.
|
||||
"""
|
||||
prompt: (
|
||||
StringPromptTemplate
|
||||
| list[StringPromptTemplate | ImagePromptTemplate | DictPromptTemplate]
|
||||
)
|
||||
if isinstance(template, str):
|
||||
prompt: StringPromptTemplate | list = PromptTemplate.from_template(
|
||||
prompt = PromptTemplate.from_template(
|
||||
template,
|
||||
template_format=template_format,
|
||||
partial_variables=partial_variables,
|
||||
)
|
||||
return cls(prompt=prompt, **kwargs)
|
||||
if isinstance(template, list):
|
||||
if isinstance(template, Sequence):
|
||||
if (partial_variables is not None) and len(partial_variables) > 0:
|
||||
msg = "Partial variables are not supported for list of templates."
|
||||
raise ValueError(msg)
|
||||
@@ -595,18 +599,18 @@ class _StringImageMessagePromptTemplate(BaseMessagePromptTemplate):
|
||||
return self._msg_class(
|
||||
content=text, additional_kwargs=self.additional_kwargs
|
||||
)
|
||||
content: list = []
|
||||
content: list[str | dict[str, Any]] = []
|
||||
for prompt in self.prompt:
|
||||
inputs = {var: kwargs[var] for var in prompt.input_variables}
|
||||
if isinstance(prompt, StringPromptTemplate):
|
||||
formatted_text: str = prompt.format(**inputs)
|
||||
formatted_text = prompt.format(**inputs)
|
||||
if formatted_text != "":
|
||||
content.append({"type": "text", "text": formatted_text})
|
||||
elif isinstance(prompt, ImagePromptTemplate):
|
||||
formatted_image: ImageURL = prompt.format(**inputs)
|
||||
formatted_image = prompt.format(**inputs)
|
||||
content.append({"type": "image_url", "image_url": formatted_image})
|
||||
elif isinstance(prompt, DictPromptTemplate):
|
||||
formatted_dict: dict[str, Any] = prompt.format(**inputs)
|
||||
formatted_dict = prompt.format(**inputs)
|
||||
content.append(formatted_dict)
|
||||
return self._msg_class(
|
||||
content=content, additional_kwargs=self.additional_kwargs
|
||||
@@ -626,18 +630,18 @@ class _StringImageMessagePromptTemplate(BaseMessagePromptTemplate):
|
||||
return self._msg_class(
|
||||
content=text, additional_kwargs=self.additional_kwargs
|
||||
)
|
||||
content: list = []
|
||||
content: list[str | dict[str, Any]] = []
|
||||
for prompt in self.prompt:
|
||||
inputs = {var: kwargs[var] for var in prompt.input_variables}
|
||||
if isinstance(prompt, StringPromptTemplate):
|
||||
formatted_text: str = await prompt.aformat(**inputs)
|
||||
formatted_text = await prompt.aformat(**inputs)
|
||||
if formatted_text != "":
|
||||
content.append({"type": "text", "text": formatted_text})
|
||||
elif isinstance(prompt, ImagePromptTemplate):
|
||||
formatted_image: ImageURL = await prompt.aformat(**inputs)
|
||||
formatted_image = await prompt.aformat(**inputs)
|
||||
content.append({"type": "image_url", "image_url": formatted_image})
|
||||
elif isinstance(prompt, DictPromptTemplate):
|
||||
formatted_dict: dict[str, Any] = prompt.format(**inputs)
|
||||
formatted_dict = prompt.format(**inputs)
|
||||
content.append(formatted_dict)
|
||||
return self._msg_class(
|
||||
content=content, additional_kwargs=self.additional_kwargs
|
||||
@@ -688,12 +692,12 @@ class SystemMessagePromptTemplate(_StringImageMessagePromptTemplate):
|
||||
_msg_class: type[BaseMessage] = SystemMessage
|
||||
|
||||
|
||||
class BaseChatPromptTemplate(BasePromptTemplate, ABC):
|
||||
class BaseChatPromptTemplate(BasePromptTemplate[str], ABC):
|
||||
"""Base class for chat prompt templates."""
|
||||
|
||||
@property
|
||||
@override
|
||||
def lc_attributes(self) -> dict:
|
||||
def lc_attributes(self) -> dict[str, Any]:
|
||||
return {"input_variables": self.input_variables}
|
||||
|
||||
def format(self, **kwargs: Any) -> str:
|
||||
@@ -781,7 +785,7 @@ MessageLike = BaseMessagePromptTemplate | BaseMessage | BaseChatPromptTemplate
|
||||
|
||||
MessageLikeRepresentation = (
|
||||
MessageLike
|
||||
| tuple[str | type, str | Sequence[dict] | Sequence[object]]
|
||||
| tuple[str | type, str | Sequence[dict[str, Any]] | Sequence[object]]
|
||||
| str
|
||||
| dict[str, Any]
|
||||
)
|
||||
@@ -1046,7 +1050,7 @@ class ChatPromptTemplate(BaseChatPromptTemplate):
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def validate_input_variables(cls, values: dict) -> Any:
|
||||
def validate_input_variables(cls, values: dict[str, Any]) -> Any:
|
||||
"""Validate input variables.
|
||||
|
||||
If `input_variables` is not set, it will be set to the union of all input
|
||||
@@ -1062,7 +1066,7 @@ class ChatPromptTemplate(BaseChatPromptTemplate):
|
||||
ValueError: If input variables do not match.
|
||||
"""
|
||||
messages = values["messages"]
|
||||
input_vars: set = set()
|
||||
input_vars: set[str] = set()
|
||||
optional_variables = set()
|
||||
input_types: dict[str, Any] = values.get("input_types", {})
|
||||
for message in messages:
|
||||
@@ -1336,7 +1340,7 @@ class ChatPromptTemplate(BaseChatPromptTemplate):
|
||||
|
||||
def _create_template_from_message_type(
|
||||
message_type: str,
|
||||
template: str | list,
|
||||
template: str | list[str | dict[str, Any]],
|
||||
template_format: PromptTemplateFormat = "f-string",
|
||||
) -> BaseMessagePromptTemplate:
|
||||
"""Create a message prompt template from a message type and template string.
|
||||
|
||||
@@ -16,7 +16,7 @@ from langchain_core.runnables import RunnableConfig, RunnableSerializable
|
||||
from langchain_core.runnables.config import ensure_config
|
||||
|
||||
|
||||
class DictPromptTemplate(RunnableSerializable[dict, dict]):
|
||||
class DictPromptTemplate(RunnableSerializable[dict[str, Any], dict[str, Any]]):
|
||||
"""Template represented by a dictionary.
|
||||
|
||||
Recognizes variables in f-string or mustache formatted string dict values.
|
||||
@@ -74,8 +74,8 @@ class DictPromptTemplate(RunnableSerializable[dict, dict]):
|
||||
|
||||
@override
|
||||
def invoke(
|
||||
self, input: dict, config: RunnableConfig | None = None, **kwargs: Any
|
||||
) -> dict:
|
||||
self, input: dict[str, Any], config: RunnableConfig | None = None, **kwargs: Any
|
||||
) -> dict[str, Any]:
|
||||
return self._call_with_config(
|
||||
lambda x: self.format(**x),
|
||||
input,
|
||||
@@ -123,7 +123,7 @@ class DictPromptTemplate(RunnableSerializable[dict, dict]):
|
||||
|
||||
|
||||
def _get_input_variables(
|
||||
template: dict, template_format: Literal["f-string", "mustache"]
|
||||
template: dict[str, Any], template_format: Literal["f-string", "mustache"]
|
||||
) -> list[str]:
|
||||
input_variables = []
|
||||
for v in template.values():
|
||||
|
||||
@@ -34,7 +34,7 @@ if TYPE_CHECKING:
|
||||
class _FewShotPromptTemplateMixin(BaseModel):
|
||||
"""Prompt template that contains few shot examples."""
|
||||
|
||||
examples: list[dict] | None = None
|
||||
examples: list[dict[str, Any]] | None = None
|
||||
"""Examples to format into the prompt.
|
||||
|
||||
Either this or `example_selector` should be provided.
|
||||
@@ -53,7 +53,7 @@ class _FewShotPromptTemplateMixin(BaseModel):
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def check_examples_and_selector(cls, values: dict) -> Any:
|
||||
def check_examples_and_selector(cls, values: dict[str, Any]) -> Any:
|
||||
"""Check that one and only one of `examples`/`example_selector` are provided.
|
||||
|
||||
Args:
|
||||
@@ -79,7 +79,7 @@ class _FewShotPromptTemplateMixin(BaseModel):
|
||||
|
||||
return values
|
||||
|
||||
def _get_examples(self, **kwargs: Any) -> list[dict]:
|
||||
def _get_examples(self, **kwargs: Any) -> list[dict[str, Any]]:
|
||||
"""Get the examples to use for formatting the prompt.
|
||||
|
||||
Args:
|
||||
@@ -98,7 +98,7 @@ class _FewShotPromptTemplateMixin(BaseModel):
|
||||
msg = "One of 'examples' and 'example_selector' should be provided"
|
||||
raise ValueError(msg)
|
||||
|
||||
async def _aget_examples(self, **kwargs: Any) -> list[dict]:
|
||||
async def _aget_examples(self, **kwargs: Any) -> list[dict[str, Any]]:
|
||||
"""Async get the examples to use for formatting the prompt.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -19,7 +19,7 @@ from langchain_core.prompts.string import (
|
||||
class FewShotPromptWithTemplates(StringPromptTemplate):
|
||||
"""Prompt template that contains few shot examples."""
|
||||
|
||||
examples: list[dict] | None = None
|
||||
examples: list[dict[str, Any]] | None = None
|
||||
"""Examples to format into the prompt.
|
||||
|
||||
Either this or `example_selector` should be provided.
|
||||
@@ -63,7 +63,7 @@ class FewShotPromptWithTemplates(StringPromptTemplate):
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def check_examples_and_selector(cls, values: dict) -> Any:
|
||||
def check_examples_and_selector(cls, values: dict[str, Any]) -> Any:
|
||||
"""Check that one and only one of examples/example_selector are provided."""
|
||||
examples = values.get("examples")
|
||||
example_selector = values.get("example_selector")
|
||||
@@ -106,14 +106,14 @@ class FewShotPromptWithTemplates(StringPromptTemplate):
|
||||
extra="forbid",
|
||||
)
|
||||
|
||||
def _get_examples(self, **kwargs: Any) -> list[dict]:
|
||||
def _get_examples(self, **kwargs: Any) -> list[dict[str, Any]]:
|
||||
if self.examples is not None:
|
||||
return self.examples
|
||||
if self.example_selector is not None:
|
||||
return self.example_selector.select_examples(kwargs)
|
||||
raise ValueError
|
||||
|
||||
async def _aget_examples(self, **kwargs: Any) -> list[dict]:
|
||||
async def _aget_examples(self, **kwargs: Any) -> list[dict[str, Any]]:
|
||||
if self.examples is not None:
|
||||
return self.examples
|
||||
if self.example_selector is not None:
|
||||
|
||||
@@ -29,7 +29,7 @@ class ImagePromptTemplate(BasePromptTemplate[ImageURL]):
|
||||
```
|
||||
"""
|
||||
|
||||
template: dict = Field(default_factory=dict)
|
||||
template: dict[str, Any] = Field(default_factory=dict)
|
||||
"""Template for the prompt."""
|
||||
|
||||
template_format: PromptTemplateFormat = "f-string"
|
||||
|
||||
@@ -4,6 +4,7 @@ import json
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import yaml
|
||||
|
||||
@@ -52,8 +53,8 @@ def _validate_path(path: Path) -> None:
|
||||
"prompts and `load`/`loads` to deserialize them.",
|
||||
)
|
||||
def load_prompt_from_config(
|
||||
config: dict, *, allow_dangerous_paths: bool = False
|
||||
) -> BasePromptTemplate:
|
||||
config: dict[str, Any], *, allow_dangerous_paths: bool = False
|
||||
) -> BasePromptTemplate[str]:
|
||||
"""Load prompt from config dict.
|
||||
|
||||
Args:
|
||||
@@ -83,8 +84,8 @@ def load_prompt_from_config(
|
||||
|
||||
|
||||
def _load_template(
|
||||
var_name: str, config: dict, *, allow_dangerous_paths: bool = False
|
||||
) -> dict:
|
||||
var_name: str, config: dict[str, Any], *, allow_dangerous_paths: bool = False
|
||||
) -> dict[str, Any]:
|
||||
"""Load template from the path if applicable."""
|
||||
# Check if template_path exists in config.
|
||||
if f"{var_name}_path" in config:
|
||||
@@ -109,7 +110,9 @@ def _load_template(
|
||||
return config
|
||||
|
||||
|
||||
def _load_examples(config: dict, *, allow_dangerous_paths: bool = False) -> dict:
|
||||
def _load_examples(
|
||||
config: dict[str, Any], *, allow_dangerous_paths: bool = False
|
||||
) -> dict[str, Any]:
|
||||
"""Load examples if necessary."""
|
||||
if isinstance(config["examples"], list):
|
||||
pass
|
||||
@@ -132,7 +135,7 @@ def _load_examples(config: dict, *, allow_dangerous_paths: bool = False) -> dict
|
||||
return config
|
||||
|
||||
|
||||
def _load_output_parser(config: dict) -> dict:
|
||||
def _load_output_parser(config: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Load output parser."""
|
||||
if config_ := config.get("output_parser"):
|
||||
if output_parser_type := config_.get("_type") != "default":
|
||||
@@ -143,7 +146,7 @@ def _load_output_parser(config: dict) -> dict:
|
||||
|
||||
|
||||
def _load_few_shot_prompt(
|
||||
config: dict, *, allow_dangerous_paths: bool = False
|
||||
config: dict[str, Any], *, allow_dangerous_paths: bool = False
|
||||
) -> FewShotPromptTemplate:
|
||||
"""Load the "few shot" prompt from the config."""
|
||||
# Load the suffix and prefix templates.
|
||||
@@ -178,7 +181,7 @@ def _load_few_shot_prompt(
|
||||
|
||||
|
||||
def _load_prompt(
|
||||
config: dict, *, allow_dangerous_paths: bool = False
|
||||
config: dict[str, Any], *, allow_dangerous_paths: bool = False
|
||||
) -> PromptTemplate:
|
||||
"""Load the prompt template from config."""
|
||||
# Load the template from disk if necessary.
|
||||
@@ -212,7 +215,7 @@ def load_prompt(
|
||||
encoding: str | None = None,
|
||||
*,
|
||||
allow_dangerous_paths: bool = False,
|
||||
) -> BasePromptTemplate:
|
||||
) -> BasePromptTemplate[str]:
|
||||
"""Unified method for loading a prompt from LangChainHub or local filesystem.
|
||||
|
||||
Args:
|
||||
@@ -247,7 +250,7 @@ def _load_prompt_from_file(
|
||||
encoding: str | None = None,
|
||||
*,
|
||||
allow_dangerous_paths: bool = False,
|
||||
) -> BasePromptTemplate:
|
||||
) -> BasePromptTemplate[str]:
|
||||
"""Load prompt from file."""
|
||||
# Convert file to a Path object.
|
||||
file_path = Path(file)
|
||||
@@ -266,7 +269,7 @@ def _load_prompt_from_file(
|
||||
|
||||
|
||||
def _load_chat_prompt(
|
||||
config: dict,
|
||||
config: dict[str, Any],
|
||||
*,
|
||||
allow_dangerous_paths: bool = False, # noqa: ARG001
|
||||
) -> ChatPromptTemplate:
|
||||
@@ -282,7 +285,7 @@ def _load_chat_prompt(
|
||||
return ChatPromptTemplate.from_template(template=template, **config)
|
||||
|
||||
|
||||
type_to_loader_dict: dict[str, Callable[..., BasePromptTemplate]] = {
|
||||
type_to_loader_dict: dict[str, Callable[..., BasePromptTemplate[str]]] = {
|
||||
"prompt": _load_prompt,
|
||||
"few_shot": _load_few_shot_prompt,
|
||||
"chat": _load_chat_prompt,
|
||||
|
||||
@@ -88,7 +88,7 @@ class PromptTemplate(StringPromptTemplate):
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def pre_init_validation(cls, values: dict) -> Any:
|
||||
def pre_init_validation(cls, values: dict[str, Any]) -> Any:
|
||||
"""Check that template and input variables are consistent."""
|
||||
if values.get("template") is None:
|
||||
# Will let pydantic fail with a ValidationError if template
|
||||
|
||||
@@ -213,7 +213,7 @@ DEFAULT_FORMATTER_MAPPING: dict[str, Callable[..., str]] = {
|
||||
"jinja2": jinja2_formatter,
|
||||
}
|
||||
|
||||
DEFAULT_VALIDATOR_MAPPING: dict[str, Callable] = {
|
||||
DEFAULT_VALIDATOR_MAPPING: dict[str, Callable[[str, list[str]], None]] = {
|
||||
"f-string": formatter.validate_input_variables,
|
||||
"jinja2": validate_jinja2,
|
||||
}
|
||||
@@ -325,7 +325,7 @@ def get_template_variables(template: str, template_format: str) -> list[str]:
|
||||
return sorted(input_variables)
|
||||
|
||||
|
||||
class StringPromptTemplate(BasePromptTemplate, ABC):
|
||||
class StringPromptTemplate(BasePromptTemplate[str], ABC):
|
||||
"""String prompt that exposes the format method, returning a prompt."""
|
||||
|
||||
@classmethod
|
||||
@@ -390,7 +390,7 @@ class StringPromptTemplate(BasePromptTemplate, ABC):
|
||||
print(self.pretty_repr(html=is_interactive_env())) # noqa: T201
|
||||
|
||||
|
||||
def is_subsequence(child: Sequence, parent: Sequence) -> bool:
|
||||
def is_subsequence(child: Sequence[Any], parent: Sequence[Any]) -> bool:
|
||||
"""Return `True` if child is subsequence of parent."""
|
||||
if len(child) == 0 or len(parent) == 0:
|
||||
return False
|
||||
|
||||
@@ -28,7 +28,7 @@ from langchain_core.utils import get_pydantic_field_names
|
||||
class StructuredPrompt(ChatPromptTemplate):
|
||||
"""Structured prompt template for a language model."""
|
||||
|
||||
schema_: dict | type
|
||||
schema_: dict[str, Any] | type
|
||||
"""Schema for the structured prompt."""
|
||||
|
||||
structured_output_kwargs: dict[str, Any] = Field(default_factory=dict)
|
||||
@@ -36,7 +36,7 @@ class StructuredPrompt(ChatPromptTemplate):
|
||||
def __init__(
|
||||
self,
|
||||
messages: Sequence[MessageLikeRepresentation],
|
||||
schema_: dict | type[BaseModel] | None = None,
|
||||
schema_: dict[str, Any] | type[BaseModel] | None = None,
|
||||
*,
|
||||
structured_output_kwargs: dict[str, Any] | None = None,
|
||||
template_format: PromptTemplateFormat = "f-string",
|
||||
@@ -87,7 +87,7 @@ class StructuredPrompt(ChatPromptTemplate):
|
||||
def from_messages_and_schema(
|
||||
cls,
|
||||
messages: Sequence[MessageLikeRepresentation],
|
||||
schema: dict | type,
|
||||
schema: dict[str, Any] | type,
|
||||
**kwargs: Any,
|
||||
) -> ChatPromptTemplate:
|
||||
"""Create a chat prompt template from a variety of message formats.
|
||||
@@ -143,7 +143,7 @@ class StructuredPrompt(ChatPromptTemplate):
|
||||
| Callable[[AsyncIterator[Any]], AsyncIterator[Other]]
|
||||
| Callable[[Any], Other]
|
||||
| Mapping[str, Runnable[Any, Other] | Callable[[Any], Other] | Any],
|
||||
) -> RunnableSerializable[dict, Other]:
|
||||
) -> RunnableSerializable[dict[str, Any], Other]:
|
||||
return self.pipe(other)
|
||||
|
||||
def pipe(
|
||||
@@ -154,7 +154,7 @@ class StructuredPrompt(ChatPromptTemplate):
|
||||
| Callable[[Any], Other]
|
||||
| Mapping[str, Runnable[Any, Other] | Callable[[Any], Other] | Any],
|
||||
name: str | None = None,
|
||||
) -> RunnableSerializable[dict, Other]:
|
||||
) -> RunnableSerializable[dict[str, Any], Other]:
|
||||
"""Pipe the structured prompt to a language model.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -605,7 +605,7 @@ class Runnable(ABC, Generic[Input, Output]):
|
||||
|
||||
def get_prompts(
|
||||
self, config: RunnableConfig | None = None
|
||||
) -> list[BasePromptTemplate]:
|
||||
) -> list[BasePromptTemplate[Any]]:
|
||||
"""Return a list of prompts used by this `Runnable`."""
|
||||
# Import locally to prevent circular import
|
||||
from langchain_core.prompts.base import BasePromptTemplate # noqa: PLC0415
|
||||
@@ -2488,7 +2488,7 @@ class Runnable(ABC, Generic[Input, Output]):
|
||||
iterator = context.run(transformer, input_for_transform, **kwargs) # type: ignore[arg-type]
|
||||
if stream_handler := next(
|
||||
(
|
||||
cast("_StreamingCallbackHandler", h)
|
||||
h
|
||||
for h in run_manager.handlers
|
||||
# instance check OK here, it's a mixin
|
||||
if isinstance(h, _StreamingCallbackHandler)
|
||||
@@ -2590,7 +2590,7 @@ class Runnable(ABC, Generic[Input, Output]):
|
||||
|
||||
if stream_handler := next(
|
||||
(
|
||||
cast("_StreamingCallbackHandler", h)
|
||||
h
|
||||
for h in run_manager.handlers
|
||||
# instance check OK here, it's a mixin
|
||||
if isinstance(h, _StreamingCallbackHandler)
|
||||
@@ -3088,7 +3088,7 @@ class RunnableSequence(RunnableSerializable[Input, Output]):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*steps: RunnableLike,
|
||||
*steps: RunnableLike[Any, Any],
|
||||
name: str | None = None,
|
||||
first: Runnable[Any, Any] | None = None,
|
||||
middle: list[Runnable[Any, Any]] | None = None,
|
||||
@@ -3106,7 +3106,7 @@ class RunnableSequence(RunnableSerializable[Input, Output]):
|
||||
Raises:
|
||||
ValueError: If the sequence has less than 2 steps.
|
||||
"""
|
||||
steps_flat: list[Runnable] = []
|
||||
steps_flat: list[Runnable[Any, Any]] = []
|
||||
if not steps and first is not None and last is not None:
|
||||
steps_flat = [first] + (middle or []) + [last]
|
||||
for step in steps:
|
||||
@@ -4216,7 +4216,7 @@ class RunnableParallel(RunnableSerializable[Input, dict[str, Any]]):
|
||||
]
|
||||
|
||||
# Wrap in a coroutine to satisfy linter
|
||||
async def get_next_chunk(generator: AsyncIterator) -> Output | None:
|
||||
async def get_next_chunk(generator: AsyncIterator[Any]) -> Output | None:
|
||||
return await anext(generator)
|
||||
|
||||
# Start the first iteration of each generator
|
||||
@@ -4383,9 +4383,10 @@ class RunnableGenerator(Runnable[Input, Output]):
|
||||
TypeError: If the transform is not a generator function.
|
||||
|
||||
"""
|
||||
func_for_name: Callable[..., Any]
|
||||
if atransform is not None:
|
||||
self._atransform = atransform
|
||||
func_for_name: Callable = atransform
|
||||
func_for_name = atransform
|
||||
|
||||
if is_async_generator(transform):
|
||||
self._atransform = transform
|
||||
@@ -4795,9 +4796,10 @@ class RunnableLambda(Runnable[Input, Output]):
|
||||
TypeError: If both `func` and `afunc` are provided.
|
||||
|
||||
"""
|
||||
func_for_name: Callable[..., Any]
|
||||
if afunc is not None:
|
||||
self.afunc = afunc
|
||||
func_for_name: Callable = afunc
|
||||
func_for_name = afunc
|
||||
|
||||
if is_async_callable(func) or is_async_generator(func):
|
||||
if afunc is not None:
|
||||
@@ -4939,7 +4941,7 @@ class RunnableLambda(Runnable[Input, Output]):
|
||||
)
|
||||
|
||||
@functools.cached_property
|
||||
def deps(self) -> list[Runnable]:
|
||||
def deps(self) -> list[Runnable[Any, Any]]:
|
||||
"""The dependencies of this `Runnable`.
|
||||
|
||||
Returns:
|
||||
@@ -4954,7 +4956,7 @@ class RunnableLambda(Runnable[Input, Output]):
|
||||
else:
|
||||
objects = []
|
||||
|
||||
deps: list[Runnable] = []
|
||||
deps: list[Runnable[Any, Any]] = []
|
||||
for obj in objects:
|
||||
if isinstance(obj, Runnable):
|
||||
deps.append(obj)
|
||||
@@ -5130,7 +5132,7 @@ class RunnableLambda(Runnable[Input, Output]):
|
||||
cast(
|
||||
"AsyncGenerator[Any, Any]",
|
||||
acall_func_with_variable_args(
|
||||
cast("Callable", afunc),
|
||||
cast("Callable[..., Any]", afunc),
|
||||
value,
|
||||
config,
|
||||
run_manager,
|
||||
@@ -5151,7 +5153,7 @@ class RunnableLambda(Runnable[Input, Output]):
|
||||
output = chunk
|
||||
else:
|
||||
output = await acall_func_with_variable_args(
|
||||
cast("Callable", afunc), value, config, run_manager, **kwargs
|
||||
cast("Callable[..., Any]", afunc), value, config, run_manager, **kwargs
|
||||
)
|
||||
# If the output is a Runnable, invoke it
|
||||
if isinstance(output, Runnable):
|
||||
@@ -5373,7 +5375,7 @@ class RunnableLambda(Runnable[Input, Output]):
|
||||
async for chunk in cast(
|
||||
"AsyncIterator[Output]",
|
||||
acall_func_with_variable_args(
|
||||
cast("Callable", afunc),
|
||||
cast("Callable[..., Any]", afunc),
|
||||
final,
|
||||
config,
|
||||
run_manager,
|
||||
@@ -5390,7 +5392,7 @@ class RunnableLambda(Runnable[Input, Output]):
|
||||
output = chunk
|
||||
else:
|
||||
output = await acall_func_with_variable_args(
|
||||
cast("Callable", afunc),
|
||||
cast("Callable[..., Any]", afunc),
|
||||
final,
|
||||
config,
|
||||
run_manager,
|
||||
@@ -6486,7 +6488,7 @@ RunnableLike = (
|
||||
)
|
||||
|
||||
|
||||
def coerce_to_runnable(thing: RunnableLike) -> Runnable[Input, Output]:
|
||||
def coerce_to_runnable(thing: RunnableLike[Input, Output]) -> Runnable[Input, Output]:
|
||||
"""Coerce a `Runnable`-like object into a `Runnable`.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -77,9 +77,9 @@ class RunnableBranch(RunnableSerializable[Input, Output]):
|
||||
Runnable[Input, bool]
|
||||
| Callable[[Input], bool]
|
||||
| Callable[[Input], Awaitable[bool]],
|
||||
RunnableLike,
|
||||
RunnableLike[Input, Output],
|
||||
]
|
||||
| RunnableLike,
|
||||
| RunnableLike[Input, Output],
|
||||
) -> None:
|
||||
"""A `Runnable` that runs one of two branches based on a condition.
|
||||
|
||||
@@ -106,9 +106,7 @@ class RunnableBranch(RunnableSerializable[Input, Output]):
|
||||
msg = "RunnableBranch default must be Runnable, callable or mapping."
|
||||
raise TypeError(msg)
|
||||
|
||||
default_ = cast(
|
||||
"Runnable[Input, Output]", coerce_to_runnable(cast("RunnableLike", default))
|
||||
)
|
||||
default_ = coerce_to_runnable(cast("RunnableLike[Input, Output]", default))
|
||||
|
||||
branches_ = []
|
||||
|
||||
|
||||
@@ -97,7 +97,7 @@ class Node(NamedTuple):
|
||||
"""The unique identifier of the node."""
|
||||
name: str
|
||||
"""The name of the node."""
|
||||
data: type[BaseModel] | RunnableType | None
|
||||
data: type[BaseModel] | RunnableType[Any, Any] | None
|
||||
"""The data of the node."""
|
||||
metadata: dict[str, Any] | None
|
||||
"""Optional metadata for the node. """
|
||||
@@ -177,7 +177,7 @@ class MermaidDrawMethod(Enum):
|
||||
|
||||
def node_data_str(
|
||||
id: str,
|
||||
data: type[BaseModel] | RunnableType | None,
|
||||
data: type[BaseModel] | RunnableType[Any, Any] | None,
|
||||
) -> str:
|
||||
"""Convert the data of a node to a string.
|
||||
|
||||
@@ -311,7 +311,7 @@ class Graph:
|
||||
|
||||
def add_node(
|
||||
self,
|
||||
data: type[BaseModel] | RunnableType | None,
|
||||
data: type[BaseModel] | RunnableType[Any, Any] | None,
|
||||
id: str | None = None,
|
||||
*,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
|
||||
@@ -36,7 +36,7 @@ MessagesOrDictWithMessages = Sequence["BaseMessage"] | dict[str, Any]
|
||||
GetSessionHistoryCallable = Callable[..., BaseChatMessageHistory]
|
||||
|
||||
|
||||
class RunnableWithMessageHistory(RunnableBindingBase): # type: ignore[no-redef]
|
||||
class RunnableWithMessageHistory(RunnableBindingBase[Any, Any]): # type: ignore[no-redef]
|
||||
"""`Runnable` that manages chat message history for another `Runnable`.
|
||||
|
||||
A chat message history is a sequence of messages that represent a conversation.
|
||||
@@ -391,7 +391,7 @@ class RunnableWithMessageHistory(RunnableBindingBase): # type: ignore[no-redef]
|
||||
|
||||
@override
|
||||
def get_input_schema(self, config: RunnableConfig | None = None) -> type[BaseModel]:
|
||||
fields: dict = {}
|
||||
fields: dict[str, Any] = {}
|
||||
if self.input_messages_key and self.history_messages_key:
|
||||
fields[self.input_messages_key] = (
|
||||
str | BaseMessage | Sequence[BaseMessage],
|
||||
@@ -450,7 +450,7 @@ class RunnableWithMessageHistory(RunnableBindingBase): # type: ignore[no-redef]
|
||||
)
|
||||
|
||||
def _get_input_messages(
|
||||
self, input_val: str | BaseMessage | Sequence[BaseMessage] | dict
|
||||
self, input_val: str | BaseMessage | Sequence[BaseMessage] | dict[str, Any]
|
||||
) -> list[BaseMessage]:
|
||||
# If dictionary, try to pluck the single key representing messages
|
||||
if isinstance(input_val, dict):
|
||||
@@ -488,7 +488,7 @@ class RunnableWithMessageHistory(RunnableBindingBase): # type: ignore[no-redef]
|
||||
raise ValueError(msg)
|
||||
|
||||
def _get_output_messages(
|
||||
self, output_val: str | BaseMessage | Sequence[BaseMessage] | dict
|
||||
self, output_val: str | BaseMessage | Sequence[BaseMessage] | dict[str, Any]
|
||||
) -> list[BaseMessage]:
|
||||
# If dictionary, try to pluck the single key representing messages
|
||||
if isinstance(output_val, dict):
|
||||
|
||||
@@ -9,7 +9,6 @@ from collections.abc import Awaitable, Callable
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
cast,
|
||||
)
|
||||
|
||||
from pydantic import BaseModel, RootModel
|
||||
@@ -346,7 +345,7 @@ class RunnablePassthrough(RunnableSerializable[Other, Other]):
|
||||
yield chunk
|
||||
|
||||
|
||||
_graph_passthrough: RunnablePassthrough = RunnablePassthrough()
|
||||
_graph_passthrough = RunnablePassthrough[Any]()
|
||||
|
||||
|
||||
class RunnableAssign(RunnableSerializable[dict[str, Any], dict[str, Any]]):
|
||||
@@ -577,9 +576,11 @@ class RunnableAssign(RunnableSerializable[dict[str, Any], dict[str, Any]]):
|
||||
if filtered:
|
||||
yield filtered
|
||||
# yield map output
|
||||
yield cast("dict[str, Any]", first_map_chunk_future.result())
|
||||
for chunk in map_output:
|
||||
yield chunk
|
||||
first_chunk = first_map_chunk_future.result()
|
||||
if first_chunk is not None:
|
||||
yield first_chunk
|
||||
for chunk in map_output:
|
||||
yield chunk
|
||||
|
||||
@override
|
||||
def transform(
|
||||
@@ -613,7 +614,7 @@ class RunnableAssign(RunnableSerializable[dict[str, Any], dict[str, Any]]):
|
||||
**kwargs,
|
||||
)
|
||||
# start map output stream
|
||||
first_map_chunk_task: asyncio.Task = asyncio.create_task(
|
||||
first_map_chunk_task = asyncio.create_task(
|
||||
anext(map_output, None),
|
||||
)
|
||||
# consume passthrough stream
|
||||
@@ -629,9 +630,11 @@ class RunnableAssign(RunnableSerializable[dict[str, Any], dict[str, Any]]):
|
||||
if filtered:
|
||||
yield filtered
|
||||
# yield map output
|
||||
yield await first_map_chunk_task
|
||||
async for chunk in map_output:
|
||||
yield chunk
|
||||
first_chunk = await first_map_chunk_task
|
||||
if first_chunk is not None:
|
||||
yield first_chunk
|
||||
async for chunk in map_output:
|
||||
yield chunk
|
||||
|
||||
@override
|
||||
async def atransform(
|
||||
|
||||
@@ -46,7 +46,9 @@ Input = TypeVar("Input", contravariant=True) # noqa: PLC0105
|
||||
Output = TypeVar("Output", covariant=True) # noqa: PLC0105
|
||||
|
||||
|
||||
async def gated_coro(semaphore: asyncio.Semaphore, coro: Coroutine) -> Any:
|
||||
async def gated_coro(
|
||||
semaphore: asyncio.Semaphore, coro: Coroutine[Any, Any, Any]
|
||||
) -> Any:
|
||||
"""Run a coroutine with a semaphore.
|
||||
|
||||
Args:
|
||||
@@ -60,7 +62,9 @@ async def gated_coro(semaphore: asyncio.Semaphore, coro: Coroutine) -> Any:
|
||||
return await coro
|
||||
|
||||
|
||||
async def gather_with_concurrency(n: int | None, *coros: Coroutine) -> list:
|
||||
async def gather_with_concurrency(
|
||||
n: int | None, *coros: Coroutine[Any, Any, Any]
|
||||
) -> list[Any]:
|
||||
"""Gather coroutines with a limit on the number of concurrent coroutines.
|
||||
|
||||
Args:
|
||||
@@ -362,7 +366,7 @@ class GetLambdaSource(ast.NodeVisitor):
|
||||
self.source = ast.unparse(node)
|
||||
|
||||
|
||||
def get_function_first_arg_dict_keys(func: Callable) -> list[str] | None:
|
||||
def get_function_first_arg_dict_keys(func: Callable[..., Any]) -> list[str] | None:
|
||||
"""Get the keys of the first argument of a function if it is a dict.
|
||||
|
||||
Args:
|
||||
@@ -381,7 +385,7 @@ def get_function_first_arg_dict_keys(func: Callable) -> list[str] | None:
|
||||
return None
|
||||
|
||||
|
||||
def get_lambda_source(func: Callable) -> str | None:
|
||||
def get_lambda_source(func: Callable[..., Any]) -> str | None:
|
||||
"""Get the source code of a lambda function.
|
||||
|
||||
Args:
|
||||
@@ -405,7 +409,7 @@ def get_lambda_source(func: Callable) -> str | None:
|
||||
|
||||
|
||||
@lru_cache(maxsize=256)
|
||||
def get_function_nonlocals(func: Callable) -> list[Any]:
|
||||
def get_function_nonlocals(func: Callable[..., Any]) -> list[Any]:
|
||||
"""Get the nonlocal variables accessed by a function.
|
||||
|
||||
Args:
|
||||
@@ -747,7 +751,7 @@ class _RootEventFilter:
|
||||
|
||||
def is_async_generator(
|
||||
func: Any,
|
||||
) -> TypeGuard[Callable[..., AsyncIterator]]:
|
||||
) -> TypeGuard[Callable[..., AsyncIterator[Any]]]:
|
||||
"""Check if a function is an async generator.
|
||||
|
||||
Args:
|
||||
@@ -764,7 +768,7 @@ def is_async_generator(
|
||||
|
||||
def is_async_callable(
|
||||
func: Any,
|
||||
) -> TypeGuard[Callable[..., Awaitable]]:
|
||||
) -> TypeGuard[Callable[..., Awaitable[Any]]]:
|
||||
"""Check if a function is async.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -122,38 +122,12 @@ def _get_annotation_description(arg_type: type) -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def _get_filtered_args(
|
||||
inferred_model: type[BaseModel],
|
||||
func: Callable,
|
||||
*,
|
||||
filter_args: Sequence[str],
|
||||
include_injected: bool = True,
|
||||
) -> dict:
|
||||
"""Get filtered arguments from a function's signature.
|
||||
|
||||
Args:
|
||||
inferred_model: The Pydantic model inferred from the function.
|
||||
func: The function to extract arguments from.
|
||||
filter_args: Arguments to exclude from the result.
|
||||
include_injected: Whether to include injected arguments.
|
||||
|
||||
Returns:
|
||||
Dictionary of filtered arguments with their schema definitions.
|
||||
"""
|
||||
schema = inferred_model.model_json_schema()["properties"]
|
||||
valid_keys = signature(func).parameters
|
||||
return {
|
||||
k: schema[k]
|
||||
for i, (k, param) in enumerate(valid_keys.items())
|
||||
if k not in filter_args
|
||||
and (i > 0 or param.name not in {"self", "cls"})
|
||||
and (include_injected or not _is_injected_arg_type(param.annotation))
|
||||
}
|
||||
|
||||
|
||||
def _parse_python_function_docstring(
|
||||
function: Callable, annotations: dict, *, error_on_invalid_docstring: bool = False
|
||||
) -> tuple[str, dict]:
|
||||
function: Callable[..., Any],
|
||||
annotations: dict[str, Any],
|
||||
*,
|
||||
error_on_invalid_docstring: bool = False,
|
||||
) -> tuple[str, dict[str, str]]:
|
||||
"""Parse function and argument descriptions from a docstring.
|
||||
|
||||
Assumes the function docstring follows Google Python style guide.
|
||||
@@ -175,7 +149,7 @@ def _parse_python_function_docstring(
|
||||
|
||||
|
||||
def _validate_docstring_args_against_annotations(
|
||||
arg_descriptions: dict, annotations: dict
|
||||
arg_descriptions: dict[str, str], annotations: dict[str, Any]
|
||||
) -> None:
|
||||
"""Validate that docstring arguments match function annotations.
|
||||
|
||||
@@ -193,11 +167,11 @@ def _validate_docstring_args_against_annotations(
|
||||
|
||||
|
||||
def _infer_arg_descriptions(
|
||||
fn: Callable,
|
||||
fn: Callable[..., Any],
|
||||
*,
|
||||
parse_docstring: bool = False,
|
||||
error_on_invalid_docstring: bool = False,
|
||||
) -> tuple[str, dict]:
|
||||
) -> tuple[str, dict[str, str]]:
|
||||
"""Infer argument descriptions from function docstring and annotations.
|
||||
|
||||
Args:
|
||||
@@ -244,7 +218,7 @@ def _is_pydantic_annotation(annotation: Any, pydantic_version: str = "v2") -> bo
|
||||
|
||||
|
||||
def _function_annotations_are_pydantic_v1(
|
||||
signature: inspect.Signature, func: Callable
|
||||
signature: inspect.Signature, func: Callable[..., Any]
|
||||
) -> bool:
|
||||
"""Check if all Pydantic annotations in a function are from v1.
|
||||
|
||||
@@ -287,7 +261,7 @@ class _SchemaConfig:
|
||||
|
||||
def create_schema_from_function(
|
||||
model_name: str,
|
||||
func: Callable,
|
||||
func: Callable[..., Any],
|
||||
*,
|
||||
filter_args: Sequence[str] | None = None,
|
||||
parse_docstring: bool = False,
|
||||
@@ -415,7 +389,7 @@ content is normalized to the content of a `ToolMessage` with `status="error"`.
|
||||
_EMPTY_SET: frozenset[str] = frozenset()
|
||||
|
||||
|
||||
class BaseTool(RunnableSerializable[str | dict | ToolCall, Any]):
|
||||
class BaseTool(RunnableSerializable[str | dict[str, Any] | ToolCall, Any]):
|
||||
"""Base class for all LangChain tools.
|
||||
|
||||
This abstract class defines the interface that all LangChain tools must implement.
|
||||
@@ -590,7 +564,7 @@ class ChildTool(BaseTool):
|
||||
return len(keys) == 1
|
||||
|
||||
@property
|
||||
def args(self) -> dict:
|
||||
def args(self) -> dict[str, Any]:
|
||||
"""Get the tool's input arguments schema.
|
||||
|
||||
Returns:
|
||||
@@ -606,7 +580,7 @@ class ChildTool(BaseTool):
|
||||
json_schema = input_schema
|
||||
else:
|
||||
json_schema = input_schema.model_json_schema()
|
||||
return cast("dict", json_schema["properties"])
|
||||
return cast("dict[str, Any]", json_schema["properties"])
|
||||
|
||||
@property
|
||||
def tool_call_schema(self) -> ArgsSchema:
|
||||
@@ -659,7 +633,7 @@ class ChildTool(BaseTool):
|
||||
@override
|
||||
def invoke(
|
||||
self,
|
||||
input: str | dict | ToolCall,
|
||||
input: str | dict[str, Any] | ToolCall,
|
||||
config: RunnableConfig | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
@@ -669,7 +643,7 @@ class ChildTool(BaseTool):
|
||||
@override
|
||||
async def ainvoke(
|
||||
self,
|
||||
input: str | dict | ToolCall,
|
||||
input: str | dict[str, Any] | ToolCall,
|
||||
config: RunnableConfig | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
@@ -679,7 +653,7 @@ class ChildTool(BaseTool):
|
||||
# --- Tool ---
|
||||
|
||||
def _parse_input(
|
||||
self, tool_input: str | dict, tool_call_id: str | None
|
||||
self, tool_input: str | dict[str, Any], tool_call_id: str | None
|
||||
) -> str | dict[str, Any]:
|
||||
"""Parse and validate tool input using the args schema.
|
||||
|
||||
@@ -825,7 +799,7 @@ class ChildTool(BaseTool):
|
||||
kwargs["run_manager"] = kwargs["run_manager"].get_sync()
|
||||
return await run_in_executor(None, self._run, *args, **kwargs)
|
||||
|
||||
def _filter_injected_args(self, tool_input: dict) -> dict:
|
||||
def _filter_injected_args(self, tool_input: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Filter out injected tool arguments from the input dictionary.
|
||||
|
||||
Injected arguments are those annotated with `InjectedToolArg` or its
|
||||
@@ -862,8 +836,8 @@ class ChildTool(BaseTool):
|
||||
return {k: v for k, v in tool_input.items() if k not in filtered_keys}
|
||||
|
||||
def _to_args_and_kwargs(
|
||||
self, tool_input: str | dict, tool_call_id: str | None
|
||||
) -> tuple[tuple, dict]:
|
||||
self, tool_input: str | dict[str, Any], tool_call_id: str | None
|
||||
) -> tuple[tuple[str, ...], dict[str, Any]]:
|
||||
"""Convert tool input to positional and keyword arguments.
|
||||
|
||||
Args:
|
||||
@@ -1030,7 +1004,7 @@ class ChildTool(BaseTool):
|
||||
|
||||
async def arun(
|
||||
self,
|
||||
tool_input: str | dict,
|
||||
tool_input: str | dict[str, Any],
|
||||
verbose: bool | None = None, # noqa: FBT001
|
||||
start_color: str | None = "green",
|
||||
color: str | None = "green",
|
||||
@@ -1244,10 +1218,10 @@ def _handle_tool_error(
|
||||
|
||||
|
||||
def _prep_run_args(
|
||||
value: str | dict | ToolCall,
|
||||
value: str | dict[str, Any] | ToolCall,
|
||||
config: RunnableConfig | None,
|
||||
**kwargs: Any,
|
||||
) -> tuple[str | dict, dict]:
|
||||
) -> tuple[str | dict[str, Any], dict[str, Any]]:
|
||||
"""Prepare arguments for tool execution.
|
||||
|
||||
Args:
|
||||
@@ -1259,12 +1233,13 @@ def _prep_run_args(
|
||||
A tuple of `(tool_input, run_kwargs)`.
|
||||
"""
|
||||
config = ensure_config(config)
|
||||
tool_input: str | dict[str, Any]
|
||||
if _is_tool_call(value):
|
||||
tool_call_id: str | None = cast("ToolCall", value)["id"]
|
||||
tool_input: str | dict = cast("ToolCall", value)["args"].copy()
|
||||
tool_input = cast("ToolCall", value)["args"].copy()
|
||||
else:
|
||||
tool_call_id = None
|
||||
tool_input = cast("str | dict", value)
|
||||
tool_input = cast("str | dict[str, Any]", value)
|
||||
return (
|
||||
tool_input,
|
||||
dict(
|
||||
@@ -1376,7 +1351,7 @@ def _stringify(content: Any) -> str:
|
||||
return str(content)
|
||||
|
||||
|
||||
def _get_type_hints(func: Callable) -> dict[str, type] | None:
|
||||
def _get_type_hints(func: Callable[..., Any]) -> dict[str, type] | None:
|
||||
"""Get type hints from a function, handling partial functions.
|
||||
|
||||
Args:
|
||||
@@ -1393,7 +1368,7 @@ def _get_type_hints(func: Callable) -> dict[str, type] | None:
|
||||
return None
|
||||
|
||||
|
||||
def _get_runnable_config_param(func: Callable) -> str | None:
|
||||
def _get_runnable_config_param(func: Callable[..., Any]) -> str | None:
|
||||
"""Find the parameter name for `RunnableConfig` in a function.
|
||||
|
||||
Args:
|
||||
@@ -1527,6 +1502,7 @@ def get_all_basemodel_annotations(
|
||||
Returns:
|
||||
`dict` of field names to their type annotations.
|
||||
"""
|
||||
orig_bases: tuple[type, ...]
|
||||
# cls has no subscript: cls = FooBar
|
||||
if isinstance(cls, type):
|
||||
fields = get_fields(cls)
|
||||
@@ -1540,7 +1516,7 @@ def get_all_basemodel_annotations(
|
||||
continue
|
||||
field_name = alias_map.get(name, name)
|
||||
annotations[field_name] = param.annotation
|
||||
orig_bases: tuple = getattr(cls, "__orig_bases__", ())
|
||||
orig_bases = getattr(cls, "__orig_bases__", ())
|
||||
# cls has subscript: cls = FooBar[int]
|
||||
else:
|
||||
annotations = get_all_basemodel_annotations(
|
||||
@@ -1572,7 +1548,9 @@ def get_all_basemodel_annotations(
|
||||
# parent_origin = class Baz,
|
||||
# generic_type_vars = (type vars in Baz)
|
||||
# generic_map = {type var in Baz: str}
|
||||
generic_type_vars: tuple = getattr(parent_origin, "__parameters__", ())
|
||||
generic_type_vars: tuple[TypeVar, ...] = getattr(
|
||||
parent_origin, "__parameters__", ()
|
||||
)
|
||||
generic_map = dict(zip(generic_type_vars, get_args(parent), strict=False))
|
||||
for field in getattr(parent_origin, "__annotations__", {}):
|
||||
annotations[field] = _replace_type_vars(
|
||||
|
||||
@@ -24,13 +24,13 @@ def tool(
|
||||
parse_docstring: bool = False,
|
||||
error_on_invalid_docstring: bool = True,
|
||||
extras: dict[str, Any] | None = None,
|
||||
) -> Callable[[Callable | Runnable], BaseTool]: ...
|
||||
) -> Callable[[Callable[..., Any] | Runnable[Any, Any]], BaseTool]: ...
|
||||
|
||||
|
||||
@overload
|
||||
def tool(
|
||||
name_or_callable: str,
|
||||
runnable: Runnable,
|
||||
runnable: Runnable[Any, Any],
|
||||
*,
|
||||
description: str | None = None,
|
||||
return_direct: bool = False,
|
||||
@@ -45,7 +45,7 @@ def tool(
|
||||
|
||||
@overload
|
||||
def tool(
|
||||
name_or_callable: Callable,
|
||||
name_or_callable: Callable[..., Any],
|
||||
*,
|
||||
description: str | None = None,
|
||||
return_direct: bool = False,
|
||||
@@ -70,12 +70,12 @@ def tool(
|
||||
parse_docstring: bool = False,
|
||||
error_on_invalid_docstring: bool = True,
|
||||
extras: dict[str, Any] | None = None,
|
||||
) -> Callable[[Callable | Runnable], BaseTool]: ...
|
||||
) -> Callable[[Callable[..., Any] | Runnable[Any, Any]], BaseTool]: ...
|
||||
|
||||
|
||||
def tool(
|
||||
name_or_callable: str | Callable | None = None,
|
||||
runnable: Runnable | None = None,
|
||||
name_or_callable: str | Callable[..., Any] | None = None,
|
||||
runnable: Runnable[Any, Any] | None = None,
|
||||
*args: Any,
|
||||
description: str | None = None,
|
||||
return_direct: bool = False,
|
||||
@@ -85,7 +85,7 @@ def tool(
|
||||
parse_docstring: bool = False,
|
||||
error_on_invalid_docstring: bool = True,
|
||||
extras: dict[str, Any] | None = None,
|
||||
) -> BaseTool | Callable[[Callable | Runnable], BaseTool]:
|
||||
) -> BaseTool | Callable[[Callable[..., Any] | Runnable[Any, Any]], BaseTool]:
|
||||
"""Convert Python functions and `Runnables` to LangChain tools.
|
||||
|
||||
Can be used as a decorator with or without arguments to create tools from functions.
|
||||
@@ -258,7 +258,7 @@ def tool(
|
||||
|
||||
def _create_tool_factory(
|
||||
tool_name: str,
|
||||
) -> Callable[[Callable | Runnable], BaseTool]:
|
||||
) -> Callable[[Callable[..., Any] | Runnable[Any, Any]], BaseTool]:
|
||||
"""Create a decorator that takes a callable and returns a tool.
|
||||
|
||||
Args:
|
||||
@@ -268,7 +268,9 @@ def tool(
|
||||
A function that takes a callable or `Runnable` and returns a tool.
|
||||
"""
|
||||
|
||||
def _tool_factory(dec_func: Callable | Runnable) -> BaseTool:
|
||||
def _tool_factory(
|
||||
dec_func: Callable[..., Any] | Runnable[Any, Any],
|
||||
) -> BaseTool:
|
||||
tool_description = description
|
||||
if isinstance(dec_func, Runnable):
|
||||
runnable = dec_func
|
||||
@@ -381,7 +383,7 @@ def tool(
|
||||
# @tool(parse_docstring=True)
|
||||
# def my_tool():
|
||||
# pass
|
||||
def _partial(func: Callable | Runnable) -> BaseTool:
|
||||
def _partial(func: Callable[..., Any] | Runnable[Any, Any]) -> BaseTool:
|
||||
"""Partial function that takes a `Callable` and returns a tool."""
|
||||
name_ = func.get_name() if isinstance(func, Runnable) else func.__name__
|
||||
tool_factory = _create_tool_factory(name_)
|
||||
@@ -390,14 +392,14 @@ def tool(
|
||||
return _partial
|
||||
|
||||
|
||||
def _get_description_from_runnable(runnable: Runnable) -> str:
|
||||
def _get_description_from_runnable(runnable: Runnable[Any, Any]) -> str:
|
||||
"""Generate a placeholder description of a `Runnable`."""
|
||||
input_schema = runnable.input_schema.model_json_schema()
|
||||
return f"Takes {input_schema}."
|
||||
|
||||
|
||||
def _get_schema_from_runnable_and_arg_types(
|
||||
runnable: Runnable,
|
||||
runnable: Runnable[Any, Any],
|
||||
name: str,
|
||||
arg_types: dict[str, type] | None = None,
|
||||
) -> type[BaseModel]:
|
||||
@@ -417,7 +419,7 @@ def _get_schema_from_runnable_and_arg_types(
|
||||
|
||||
|
||||
def convert_runnable_to_tool(
|
||||
runnable: Runnable,
|
||||
runnable: Runnable[Any, Any],
|
||||
args_schema: type[BaseModel] | None = None,
|
||||
*,
|
||||
name: str | None = None,
|
||||
|
||||
@@ -33,7 +33,7 @@ def create_retriever_tool(
|
||||
name: str,
|
||||
description: str,
|
||||
*,
|
||||
document_prompt: BasePromptTemplate | None = None,
|
||||
document_prompt: BasePromptTemplate[str] | None = None,
|
||||
document_separator: str = "\n\n",
|
||||
response_format: Literal["content", "content_and_artifact"] = "content",
|
||||
) -> StructuredTool:
|
||||
|
||||
@@ -44,7 +44,7 @@ class Tool(BaseTool):
|
||||
@override
|
||||
async def ainvoke(
|
||||
self,
|
||||
input: str | dict | ToolCall,
|
||||
input: str | dict[str, Any] | ToolCall,
|
||||
config: RunnableConfig | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
@@ -57,7 +57,7 @@ class Tool(BaseTool):
|
||||
# --- Tool ---
|
||||
|
||||
@property
|
||||
def args(self) -> dict:
|
||||
def args(self) -> dict[str, Any]:
|
||||
"""The tool's input arguments.
|
||||
|
||||
Returns:
|
||||
@@ -70,8 +70,8 @@ class Tool(BaseTool):
|
||||
return {"tool_input": {"type": "string"}}
|
||||
|
||||
def _to_args_and_kwargs(
|
||||
self, tool_input: str | dict, tool_call_id: str | None
|
||||
) -> tuple[tuple, dict]:
|
||||
self, tool_input: str | dict[str, Any], tool_call_id: str | None
|
||||
) -> tuple[tuple[str, ...], dict[str, Any]]:
|
||||
"""Convert tool input to Pydantic model.
|
||||
|
||||
Args:
|
||||
@@ -156,7 +156,11 @@ class Tool(BaseTool):
|
||||
|
||||
# TODO: this is for backwards compatibility, remove in future
|
||||
def __init__(
|
||||
self, name: str, func: Callable | None, description: str, **kwargs: Any
|
||||
self,
|
||||
name: str,
|
||||
func: Callable[..., Any] | None,
|
||||
description: str,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize tool."""
|
||||
super().__init__(name=name, func=func, description=description, **kwargs)
|
||||
@@ -164,7 +168,7 @@ class Tool(BaseTool):
|
||||
@classmethod
|
||||
def from_function(
|
||||
cls,
|
||||
func: Callable | None,
|
||||
func: Callable[..., Any] | None,
|
||||
name: str, # We keep these required to support backwards compatibility
|
||||
description: str,
|
||||
return_direct: bool = False, # noqa: FBT001,FBT002
|
||||
|
||||
@@ -59,7 +59,7 @@ class StructuredTool(BaseTool):
|
||||
@override
|
||||
async def ainvoke(
|
||||
self,
|
||||
input: str | dict | ToolCall,
|
||||
input: str | dict[str, Any] | ToolCall,
|
||||
config: RunnableConfig | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
@@ -132,7 +132,7 @@ class StructuredTool(BaseTool):
|
||||
@classmethod
|
||||
def from_function(
|
||||
cls,
|
||||
func: Callable | None = None,
|
||||
func: Callable[..., Any] | None = None,
|
||||
coroutine: Callable[..., Awaitable[Any]] | None = None,
|
||||
name: str | None = None,
|
||||
description: str | None = None,
|
||||
@@ -263,7 +263,7 @@ class StructuredTool(BaseTool):
|
||||
)
|
||||
|
||||
|
||||
def _filter_schema_args(func: Callable) -> list[str]:
|
||||
def _filter_schema_args(func: Callable[..., Any]) -> list[str]:
|
||||
filter_args = list(FILTERED_ARGS)
|
||||
if config_param := _get_runnable_config_param(func):
|
||||
filter_args.append(config_param)
|
||||
|
||||
@@ -539,7 +539,7 @@ class BaseTracer(_TracerCore, BaseCallbackHandler, ABC):
|
||||
self._on_retriever_end(retrieval_run)
|
||||
return retrieval_run
|
||||
|
||||
def __deepcopy__(self, memo: dict) -> BaseTracer:
|
||||
def __deepcopy__(self, memo: dict[int, Any] | None = None) -> BaseTracer:
|
||||
"""Return self."""
|
||||
return self
|
||||
|
||||
|
||||
@@ -555,7 +555,7 @@ class _TracerCore(ABC):
|
||||
retrieval_run.events.append({"name": "error", "time": retrieval_run.end_time})
|
||||
return retrieval_run
|
||||
|
||||
def __deepcopy__(self, memo: dict) -> _TracerCore:
|
||||
def __deepcopy__(self, memo: dict[int, Any] | None = None) -> _TracerCore:
|
||||
"""Return self deepcopied."""
|
||||
return self
|
||||
|
||||
|
||||
@@ -56,7 +56,7 @@ class EvaluatorCallbackHandler(BaseTracer):
|
||||
executor: ThreadPoolExecutor | None = None
|
||||
"""The thread pool executor used for running the evaluators."""
|
||||
|
||||
futures: weakref.WeakSet[Future] = weakref.WeakSet()
|
||||
futures: weakref.WeakSet[Future[None]] = weakref.WeakSet()
|
||||
"""The set of futures representing the running evaluators."""
|
||||
|
||||
skip_unfinished: bool = True
|
||||
|
||||
@@ -98,7 +98,9 @@ def _assign_name(name: str | None, serialized: dict[str, Any] | None) -> str:
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
class _AstreamEventsCallbackHandler(AsyncCallbackHandler, _StreamingCallbackHandler):
|
||||
class _AstreamEventsCallbackHandler(
|
||||
AsyncCallbackHandler, _StreamingCallbackHandler[Any]
|
||||
):
|
||||
"""An implementation of an async callback handler for astream events."""
|
||||
|
||||
def __init__(
|
||||
@@ -500,7 +502,7 @@ class _AstreamEventsCallbackHandler(AsyncCallbackHandler, _StreamingCallbackHand
|
||||
inputs_ = run_info.get("inputs")
|
||||
|
||||
generations: list[list[GenerationChunk]] | list[list[ChatGenerationChunk]]
|
||||
output: dict | BaseMessage = {}
|
||||
output: dict[str, Any] | BaseMessage = {}
|
||||
|
||||
if run_info["run_type"] == "chat_model":
|
||||
generations = cast("list[list[ChatGenerationChunk]]", response.generations)
|
||||
@@ -815,7 +817,9 @@ class _AstreamEventsCallbackHandler(AsyncCallbackHandler, _StreamingCallbackHand
|
||||
run_info["run_type"],
|
||||
)
|
||||
|
||||
def __deepcopy__(self, memo: dict) -> _AstreamEventsCallbackHandler:
|
||||
def __deepcopy__(
|
||||
self, memo: dict[int, Any] | None = None
|
||||
) -> _AstreamEventsCallbackHandler:
|
||||
"""Return self."""
|
||||
return self
|
||||
|
||||
|
||||
@@ -229,7 +229,7 @@ class RunLog(RunLogPatch):
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
class LogStreamCallbackHandler(BaseTracer, _StreamingCallbackHandler):
|
||||
class LogStreamCallbackHandler(BaseTracer, _StreamingCallbackHandler[Any]):
|
||||
"""Tracer that streams run logs to a stream."""
|
||||
|
||||
def __init__(
|
||||
|
||||
@@ -11,14 +11,14 @@ code.
|
||||
import asyncio
|
||||
from asyncio import AbstractEventLoop, Queue
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Generic, TypeVar
|
||||
from typing import Any, Generic, TypeVar
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
class _SendStream(Generic[T]):
|
||||
def __init__(
|
||||
self, reader_loop: AbstractEventLoop, queue: Queue, done: object
|
||||
self, reader_loop: AbstractEventLoop, queue: Queue[Any], done: object
|
||||
) -> None:
|
||||
"""Create a writer for the queue and done object.
|
||||
|
||||
@@ -84,7 +84,7 @@ class _SendStream(Generic[T]):
|
||||
|
||||
|
||||
class _ReceiveStream(Generic[T]):
|
||||
def __init__(self, queue: Queue, done: object) -> None:
|
||||
def __init__(self, queue: Queue[Any], done: object) -> None:
|
||||
"""Create a reader for the queue and done object.
|
||||
|
||||
This reader should be used in the same loop as the loop that was passed to the
|
||||
@@ -126,7 +126,7 @@ class _MemoryStream(Generic[T]):
|
||||
to this constructor. This will NOT be validated at run time.
|
||||
"""
|
||||
self._loop = loop
|
||||
self._queue: asyncio.Queue = asyncio.Queue(maxsize=0)
|
||||
self._queue = asyncio.Queue[Any](maxsize=0)
|
||||
self._done = object()
|
||||
|
||||
def get_send_stream(self) -> _SendStream[T]:
|
||||
|
||||
@@ -86,7 +86,7 @@ def merge_dicts(left: dict[str, Any], *others: dict[str, Any]) -> dict[str, Any]
|
||||
return merged
|
||||
|
||||
|
||||
def merge_lists(left: list | None, *others: list | None) -> list | None:
|
||||
def merge_lists(left: list[Any] | None, *others: list[Any] | None) -> list[Any] | None:
|
||||
"""Add many lists, handling `None`.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -277,7 +277,7 @@ class Tee(Generic[T]):
|
||||
atee = Tee
|
||||
|
||||
|
||||
class aclosing(AbstractAsyncContextManager): # noqa: N801
|
||||
class aclosing(AbstractAsyncContextManager[Any]): # noqa: N801
|
||||
"""Async context manager to wrap an `AsyncGenerator` that has a `aclose()` method.
|
||||
|
||||
Code like this:
|
||||
|
||||
@@ -21,7 +21,7 @@ class StrictFormatter(Formatter):
|
||||
"""
|
||||
|
||||
def vformat(
|
||||
self, format_string: str, args: Sequence, kwargs: Mapping[str, Any]
|
||||
self, format_string: str, args: Sequence[Any], kwargs: Mapping[str, Any]
|
||||
) -> str:
|
||||
"""Format a string using only keyword arguments.
|
||||
|
||||
|
||||
@@ -72,7 +72,7 @@ class FunctionDescription(TypedDict):
|
||||
description: str
|
||||
"""A description of the function."""
|
||||
|
||||
parameters: dict
|
||||
parameters: dict[str, Any]
|
||||
"""The parameters of the function."""
|
||||
|
||||
|
||||
@@ -86,7 +86,7 @@ class ToolDescription(TypedDict):
|
||||
"""The function description."""
|
||||
|
||||
|
||||
def _rm_titles(kv: dict, prev_key: str = "") -> dict:
|
||||
def _rm_titles(kv: dict[str, Any], prev_key: str = "") -> dict[str, Any]:
|
||||
"""Recursively removes `'title'` fields from a JSON schema dictionary.
|
||||
|
||||
Remove `'title'` fields from the input JSON schema dictionary,
|
||||
@@ -121,7 +121,7 @@ def _rm_titles(kv: dict, prev_key: str = "") -> dict:
|
||||
|
||||
|
||||
def _convert_json_schema_to_openai_function(
|
||||
schema: dict,
|
||||
schema: dict[str, Any],
|
||||
*,
|
||||
name: str | None = None,
|
||||
description: str | None = None,
|
||||
@@ -206,13 +206,13 @@ def _convert_pydantic_to_openai_function(
|
||||
)
|
||||
|
||||
|
||||
def _get_python_function_name(function: Callable) -> str:
|
||||
def _get_python_function_name(function: Callable[..., Any]) -> str:
|
||||
"""Get the name of a Python function."""
|
||||
return function.__name__
|
||||
|
||||
|
||||
def _convert_python_function_to_openai_function(
|
||||
function: Callable,
|
||||
function: Callable[..., Any],
|
||||
) -> FunctionDescription:
|
||||
"""Convert a Python function to an OpenAI function-calling API compatible dict.
|
||||
|
||||
@@ -243,7 +243,7 @@ def _convert_python_function_to_openai_function(
|
||||
|
||||
|
||||
def _convert_typed_dict_to_openai_function(typed_dict: type) -> FunctionDescription:
|
||||
visited: dict = {}
|
||||
visited: dict[type, type] = {}
|
||||
|
||||
model = cast(
|
||||
"type[BaseModel]",
|
||||
@@ -279,7 +279,7 @@ def _convert_any_typed_dicts_to_pydantic(
|
||||
description, arg_descriptions = _parse_google_docstring(
|
||||
docstring, list(annotations_)
|
||||
)
|
||||
fields: dict = {}
|
||||
fields: dict[str, Any] = {}
|
||||
for arg, arg_type in annotations_.items():
|
||||
if get_origin(arg_type) in {Annotated, typing_extensions.Annotated}:
|
||||
annotated_args = get_args(arg_type)
|
||||
@@ -373,7 +373,7 @@ def _format_tool_to_openai_function(tool: BaseTool) -> FunctionDescription:
|
||||
|
||||
|
||||
def convert_to_openai_function(
|
||||
function: Mapping[str, Any] | type | Callable | BaseTool,
|
||||
function: Mapping[str, Any] | type | Callable[..., Any] | BaseTool,
|
||||
*,
|
||||
strict: bool | None = None,
|
||||
) -> dict[str, Any]:
|
||||
@@ -437,16 +437,19 @@ def convert_to_openai_function(
|
||||
if function_copy and "properties" in function_copy:
|
||||
oai_function["parameters"] = function_copy
|
||||
elif isinstance(function, type) and is_basemodel_subclass(function):
|
||||
oai_function = cast("dict", _convert_pydantic_to_openai_function(function))
|
||||
oai_function = cast(
|
||||
"dict[str, Any]", _convert_pydantic_to_openai_function(function)
|
||||
)
|
||||
elif is_typeddict(function):
|
||||
oai_function = cast(
|
||||
"dict", _convert_typed_dict_to_openai_function(cast("type", function))
|
||||
"dict[str, Any]",
|
||||
_convert_typed_dict_to_openai_function(cast("type", function)),
|
||||
)
|
||||
elif isinstance(function, langchain_core.tools.base.BaseTool):
|
||||
oai_function = cast("dict", _format_tool_to_openai_function(function))
|
||||
oai_function = cast("dict[str, Any]", _format_tool_to_openai_function(function))
|
||||
elif callable(function):
|
||||
oai_function = cast(
|
||||
"dict", _convert_python_function_to_openai_function(function)
|
||||
"dict[str, Any]", _convert_python_function_to_openai_function(function)
|
||||
)
|
||||
else:
|
||||
if isinstance(function, dict) and (
|
||||
@@ -514,7 +517,7 @@ _WellKnownOpenAITools = (
|
||||
|
||||
|
||||
def convert_to_openai_tool(
|
||||
tool: Mapping[str, Any] | type[BaseModel] | Callable | BaseTool,
|
||||
tool: Mapping[str, Any] | type[BaseModel] | Callable[..., Any] | BaseTool,
|
||||
*,
|
||||
strict: bool | None = None,
|
||||
) -> dict[str, Any]:
|
||||
@@ -576,7 +579,7 @@ def convert_to_openai_tool(
|
||||
|
||||
|
||||
def convert_to_json_schema(
|
||||
schema: dict[str, Any] | type[BaseModel] | Callable | BaseTool,
|
||||
schema: dict[str, Any] | type[BaseModel] | Callable[..., Any] | BaseTool,
|
||||
*,
|
||||
strict: bool | None = None,
|
||||
) -> dict[str, Any]:
|
||||
@@ -731,7 +734,7 @@ def _parse_google_docstring(
|
||||
args: list[str],
|
||||
*,
|
||||
error_on_invalid_docstring: bool = False,
|
||||
) -> tuple[str, dict]:
|
||||
) -> tuple[str, dict[str, str]]:
|
||||
"""Parse the function and argument descriptions from the docstring of a function.
|
||||
|
||||
Assumes the function docstring follows Google Python style guide.
|
||||
|
||||
@@ -44,7 +44,7 @@ DEFAULT_LINK_REGEX = (
|
||||
|
||||
|
||||
def find_all_links(
|
||||
raw_html: str, *, pattern: str | re.Pattern | None = None
|
||||
raw_html: str, *, pattern: str | re.Pattern[str] | None = None
|
||||
) -> list[str]:
|
||||
"""Extract all links from a raw HTML string.
|
||||
|
||||
@@ -64,7 +64,7 @@ def extract_sub_links(
|
||||
url: str,
|
||||
*,
|
||||
base_url: str | None = None,
|
||||
pattern: str | re.Pattern | None = None,
|
||||
pattern: str | re.Pattern[str] | None = None,
|
||||
prevent_outside: bool = True,
|
||||
exclude_prefixes: Sequence[str] = (),
|
||||
continue_on_failure: bool = False,
|
||||
|
||||
@@ -12,7 +12,7 @@ _TEXT_COLOR_MAPPING = {
|
||||
|
||||
|
||||
def get_color_mapping(
|
||||
items: list[str], excluded_colors: list | None = None
|
||||
items: list[str], excluded_colors: list[str] | None = None
|
||||
) -> dict[str, str]:
|
||||
"""Get mapping for items to a support color.
|
||||
|
||||
|
||||
@@ -191,7 +191,9 @@ def _parse_json(
|
||||
return parser(json_str)
|
||||
|
||||
|
||||
def parse_and_check_json_markdown(text: str, expected_keys: list[str]) -> dict:
|
||||
def parse_and_check_json_markdown(
|
||||
text: str, expected_keys: list[str]
|
||||
) -> dict[str, Any]:
|
||||
"""Parse and check a JSON string from a Markdown string.
|
||||
|
||||
Checks that it contains the expected keys.
|
||||
|
||||
@@ -9,7 +9,7 @@ if TYPE_CHECKING:
|
||||
from collections.abc import Sequence
|
||||
|
||||
|
||||
def _retrieve_ref(path: str, schema: dict) -> list | dict:
|
||||
def _retrieve_ref(path: str, schema: dict[str, Any]) -> list[Any] | dict[Any, Any]:
|
||||
"""Retrieve a referenced object from a JSON schema using a path.
|
||||
|
||||
Resolves JSON schema references (e.g., `'#/definitions/MyType'`) by traversing the
|
||||
@@ -33,7 +33,7 @@ def _retrieve_ref(path: str, schema: dict) -> list | dict:
|
||||
"with #."
|
||||
)
|
||||
raise ValueError(msg)
|
||||
out: list | dict = schema
|
||||
out: list[Any] | dict[Any, Any] = schema
|
||||
for component in components[1:]:
|
||||
if component in out:
|
||||
if isinstance(out, list):
|
||||
@@ -186,11 +186,11 @@ def _dereference_refs_helper(
|
||||
|
||||
|
||||
def dereference_refs(
|
||||
schema_obj: dict,
|
||||
schema_obj: dict[str, Any],
|
||||
*,
|
||||
full_schema: dict | None = None,
|
||||
full_schema: dict[str, Any] | None = None,
|
||||
skip_keys: Sequence[str] | None = None,
|
||||
) -> dict:
|
||||
) -> dict[str, Any]:
|
||||
"""Resolve and inline JSON Schema `$ref` references in a schema object.
|
||||
|
||||
This function processes a JSON Schema and resolves all `$ref` references by
|
||||
@@ -266,7 +266,7 @@ def dereference_refs(
|
||||
keys_to_skip = list(skip_keys) if skip_keys is not None else ["$defs"]
|
||||
shallow = skip_keys is None
|
||||
return cast(
|
||||
"dict",
|
||||
"dict[str, Any]",
|
||||
_dereference_refs_helper(
|
||||
schema_obj, full, None, keys_to_skip, shallow_refs=shallow
|
||||
),
|
||||
|
||||
@@ -126,7 +126,9 @@ def is_basemodel_instance(obj: Any) -> bool:
|
||||
|
||||
|
||||
# How to type hint this?
|
||||
def pre_init(func: Callable) -> Any:
|
||||
def pre_init(
|
||||
func: Callable[[Any, dict[str, Any]], Any],
|
||||
) -> Callable[[Any, dict[str, Any]], Any]:
|
||||
"""Decorator to run a function before model initialization.
|
||||
|
||||
Args:
|
||||
@@ -202,9 +204,9 @@ class _IgnoreUnserializable(GenerateJsonSchema):
|
||||
def _create_subset_model_v1(
|
||||
name: str,
|
||||
model: type[BaseModelV1],
|
||||
field_names: list,
|
||||
field_names: list[str],
|
||||
*,
|
||||
descriptions: dict | None = None,
|
||||
descriptions: dict[str, str] | None = None,
|
||||
fn_description: str | None = None,
|
||||
) -> type[BaseModelV1]:
|
||||
"""Create a Pydantic model with only a subset of model's fields."""
|
||||
@@ -233,7 +235,7 @@ def _create_subset_model_v2(
|
||||
model: type[BaseModel],
|
||||
field_names: list[str],
|
||||
*,
|
||||
descriptions: dict | None = None,
|
||||
descriptions: dict[str, str] | None = None,
|
||||
fn_description: str | None = None,
|
||||
) -> type[BaseModel]:
|
||||
"""Create a Pydantic model with a subset of the model fields."""
|
||||
@@ -283,7 +285,7 @@ def _create_subset_model(
|
||||
model: TypeBaseModel,
|
||||
field_names: list[str],
|
||||
*,
|
||||
descriptions: dict | None = None,
|
||||
descriptions: dict[str, str] | None = None,
|
||||
fn_description: str | None = None,
|
||||
) -> type[BaseModel]:
|
||||
"""Create subset model using the same pydantic version as the input model.
|
||||
|
||||
@@ -22,7 +22,7 @@ def stringify_value(val: Any) -> str:
|
||||
return str(val)
|
||||
|
||||
|
||||
def stringify_dict(data: dict) -> str:
|
||||
def stringify_dict(data: dict[Any, Any]) -> str:
|
||||
"""Stringify a dictionary.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -1,17 +1,18 @@
|
||||
"""Usage utilities."""
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
|
||||
def _dict_int_op(
|
||||
left: dict,
|
||||
right: dict,
|
||||
left: dict[str, Any],
|
||||
right: dict[str, Any],
|
||||
op: Callable[[int, int], int],
|
||||
*,
|
||||
default: int = 0,
|
||||
depth: int = 0,
|
||||
max_depth: int = 100,
|
||||
) -> dict:
|
||||
) -> dict[str, Any]:
|
||||
"""Apply an integer operation to corresponding values in two dictionaries.
|
||||
|
||||
Recursively combines two dictionaries by applying the given operation to integer
|
||||
@@ -36,7 +37,7 @@ def _dict_int_op(
|
||||
if depth >= max_depth:
|
||||
msg = f"{max_depth=} exceeded, unable to combine dicts."
|
||||
raise ValueError(msg)
|
||||
combined: dict = {}
|
||||
combined: dict[str, Any] = {}
|
||||
for k in set(left).union(right):
|
||||
if isinstance(left.get(k, default), int) and isinstance(
|
||||
right.get(k, default), int
|
||||
|
||||
@@ -21,7 +21,7 @@ from langchain_core.utils.pydantic import (
|
||||
)
|
||||
|
||||
|
||||
def xor_args(*arg_groups: tuple[str, ...]) -> Callable:
|
||||
def xor_args(*arg_groups: tuple[str, ...]) -> Callable[..., Any]:
|
||||
"""Validate specified keyword args are mutually exclusive.
|
||||
|
||||
Args:
|
||||
@@ -31,7 +31,7 @@ def xor_args(*arg_groups: tuple[str, ...]) -> Callable:
|
||||
Decorator that validates the specified keyword args are mutually exclusive.
|
||||
"""
|
||||
|
||||
def decorator(func: Callable) -> Callable:
|
||||
def decorator(func: Callable[..., Any]) -> Callable[..., Any]:
|
||||
@functools.wraps(func)
|
||||
def wrapper(*args: Any, **kwargs: Any) -> Any:
|
||||
"""Validate exactly one arg in each group is not None."""
|
||||
|
||||
@@ -46,7 +46,7 @@ class VectorStore(ABC):
|
||||
def add_texts(
|
||||
self,
|
||||
texts: Iterable[str],
|
||||
metadatas: list[dict] | None = None,
|
||||
metadatas: list[dict[str, Any]] | None = None,
|
||||
*,
|
||||
ids: list[str] | None = None,
|
||||
**kwargs: Any,
|
||||
@@ -185,7 +185,7 @@ class VectorStore(ABC):
|
||||
async def aadd_texts(
|
||||
self,
|
||||
texts: Iterable[str],
|
||||
metadatas: list[dict] | None = None,
|
||||
metadatas: list[dict[str, Any]] | None = None,
|
||||
*,
|
||||
ids: list[str] | None = None,
|
||||
**kwargs: Any,
|
||||
@@ -849,7 +849,7 @@ class VectorStore(ABC):
|
||||
cls: type[VST],
|
||||
texts: list[str],
|
||||
embedding: Embeddings,
|
||||
metadatas: list[dict] | None = None,
|
||||
metadatas: list[dict[str, Any]] | None = None,
|
||||
*,
|
||||
ids: list[str] | None = None,
|
||||
**kwargs: Any,
|
||||
@@ -872,7 +872,7 @@ class VectorStore(ABC):
|
||||
cls,
|
||||
texts: list[str],
|
||||
embedding: Embeddings,
|
||||
metadatas: list[dict] | None = None,
|
||||
metadatas: list[dict[str, Any]] | None = None,
|
||||
*,
|
||||
ids: list[str] | None = None,
|
||||
**kwargs: Any,
|
||||
@@ -970,7 +970,7 @@ class VectorStoreRetriever(BaseRetriever):
|
||||
search_type: str = "similarity"
|
||||
"""Type of search to perform."""
|
||||
|
||||
search_kwargs: dict = Field(default_factory=dict)
|
||||
search_kwargs: dict[str, Any] = Field(default_factory=dict)
|
||||
"""Keyword arguments to pass to the search function."""
|
||||
|
||||
allowed_search_types: ClassVar[Collection[str]] = (
|
||||
@@ -985,7 +985,7 @@ class VectorStoreRetriever(BaseRetriever):
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def validate_search_type(cls, values: dict) -> Any:
|
||||
def validate_search_type(cls, values: dict[str, Any]) -> Any:
|
||||
"""Validate search type.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -489,7 +489,7 @@ class InMemoryVectorStore(VectorStore):
|
||||
cls,
|
||||
texts: list[str],
|
||||
embedding: Embeddings,
|
||||
metadatas: list[dict] | None = None,
|
||||
metadatas: list[dict[str, Any]] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> InMemoryVectorStore:
|
||||
store = cls(
|
||||
@@ -504,7 +504,7 @@ class InMemoryVectorStore(VectorStore):
|
||||
cls,
|
||||
texts: list[str],
|
||||
embedding: Embeddings,
|
||||
metadatas: list[dict] | None = None,
|
||||
metadatas: list[dict[str, Any]] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> InMemoryVectorStore:
|
||||
store = cls(
|
||||
|
||||
@@ -10,7 +10,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import warnings
|
||||
from typing import TYPE_CHECKING, cast
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
try:
|
||||
import numpy as np
|
||||
@@ -27,12 +27,16 @@ except ImportError:
|
||||
_HAS_SIMSIMD = False
|
||||
|
||||
if TYPE_CHECKING:
|
||||
Matrix = list[list[float]] | list[np.ndarray] | np.ndarray
|
||||
import numpy.typing as npt
|
||||
|
||||
Matrix = (
|
||||
list[list[float]] | list[npt.NDArray[np.floating]] | npt.NDArray[np.floating]
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _cosine_similarity(x: Matrix, y: Matrix) -> np.ndarray:
|
||||
def _cosine_similarity(x: Matrix, y: Matrix) -> npt.NDArray[np.floating]:
|
||||
"""Row-wise cosine similarity between two equal-width matrices.
|
||||
|
||||
Args:
|
||||
@@ -91,12 +95,14 @@ def _cosine_similarity(x: Matrix, y: Matrix) -> np.ndarray:
|
||||
y_norm = np.linalg.norm(y, axis=1)
|
||||
# Ignore divide by zero errors run time warnings as those are handled below.
|
||||
with np.errstate(divide="ignore", invalid="ignore"):
|
||||
similarity = np.dot(x, y.T) / np.outer(x_norm, y_norm)
|
||||
similarity: npt.NDArray[np.floating] = np.dot(x, y.T) / np.outer(
|
||||
x_norm, y_norm
|
||||
)
|
||||
if np.isnan(similarity).all():
|
||||
msg = "NaN values found, please remove the NaN values and try again"
|
||||
raise ValueError(msg) from None
|
||||
similarity[np.isnan(similarity) | np.isinf(similarity)] = 0.0
|
||||
return cast("np.ndarray", similarity)
|
||||
return similarity
|
||||
|
||||
x = np.array(x, dtype=np.float32)
|
||||
y = np.array(y, dtype=np.float32)
|
||||
@@ -104,8 +110,8 @@ def _cosine_similarity(x: Matrix, y: Matrix) -> np.ndarray:
|
||||
|
||||
|
||||
def maximal_marginal_relevance(
|
||||
query_embedding: np.ndarray,
|
||||
embedding_list: list,
|
||||
query_embedding: npt.NDArray[np.floating],
|
||||
embedding_list: list[list[float]],
|
||||
lambda_mult: float = 0.5,
|
||||
k: int = 4,
|
||||
) -> list[int]:
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
from pathlib import Path
|
||||
|
||||
from langchain_core.documents import Document
|
||||
from langchain_core.documents.base import Blob
|
||||
|
||||
|
||||
def test_init() -> None:
|
||||
@@ -10,3 +13,15 @@ def test_init() -> None:
|
||||
Document(page_content="foo", id=1),
|
||||
]:
|
||||
assert isinstance(doc, Document)
|
||||
|
||||
|
||||
def test_metadata_allows_non_string_keys(tmp_path: Path) -> None:
|
||||
metadata = {1: "one"}
|
||||
|
||||
doc = Document(page_content="foo", metadata=metadata)
|
||||
blob_from_data = Blob.from_data("foo", metadata=metadata)
|
||||
blob_from_path = Blob.from_path(tmp_path / "foo.txt", metadata=metadata)
|
||||
|
||||
assert doc.metadata == metadata
|
||||
assert blob_from_data.metadata == metadata
|
||||
assert blob_from_path.metadata == metadata
|
||||
@@ -15,7 +15,7 @@ from langchain_core.vectorstores import VectorStore
|
||||
class DummyVectorStore(VectorStore):
|
||||
def __init__(self, init_arg: str | None = None):
|
||||
self.texts: list[str] = []
|
||||
self.metadatas: list[dict] = []
|
||||
self.metadatas: list[dict[str, Any]] = []
|
||||
self._embeddings: Embeddings | None = None
|
||||
self.init_arg = init_arg
|
||||
|
||||
@@ -27,7 +27,7 @@ class DummyVectorStore(VectorStore):
|
||||
def add_texts(
|
||||
self,
|
||||
texts: Iterable[str],
|
||||
metadatas: list[dict] | None = None,
|
||||
metadatas: list[dict[str, Any]] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> list[str]:
|
||||
self.texts.extend(texts)
|
||||
@@ -66,7 +66,7 @@ class DummyVectorStore(VectorStore):
|
||||
cls,
|
||||
texts: list[str],
|
||||
embedding: Embeddings,
|
||||
metadatas: list[dict] | None = None,
|
||||
metadatas: list[dict[str, Any]] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> "DummyVectorStore":
|
||||
store = DummyVectorStore(**kwargs)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import cast
|
||||
from typing import Any, cast
|
||||
|
||||
from langchain_core.load import dumpd, load
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk
|
||||
@@ -355,11 +355,11 @@ def test_content_blocks() -> None:
|
||||
assert chunk.content_blocks == chunk.tool_calls
|
||||
|
||||
# test v1 content
|
||||
chunk_1.content = cast("str | list[str | dict]", chunk_1.content_blocks)
|
||||
chunk_1.content = cast("list[str | dict[str, Any]]", chunk_1.content_blocks)
|
||||
assert len(chunk_1.content) == 1
|
||||
chunk_1.content[0]["extras"] = {"baz": "qux"} # type: ignore[index]
|
||||
chunk_1.response_metadata["output_version"] = "v1"
|
||||
chunk_2.content = cast("str | list[str | dict]", chunk_2.content_blocks)
|
||||
chunk_2.content = cast("list[str | dict[str, Any]]", chunk_2.content_blocks)
|
||||
|
||||
chunk = chunk_1 + chunk_2 + chunk_3
|
||||
assert chunk.content == [
|
||||
@@ -379,8 +379,12 @@ def test_content_blocks() -> None:
|
||||
standard_content_2: list[types.ContentBlock] = [
|
||||
{"type": "non_standard", "index": 0, "value": {"foo": "baz"}}
|
||||
]
|
||||
chunk_1 = AIMessageChunk(content=cast("str | list[str | dict]", standard_content_1))
|
||||
chunk_2 = AIMessageChunk(content=cast("str | list[str | dict]", standard_content_2))
|
||||
chunk_1 = AIMessageChunk(
|
||||
content=cast("list[str | dict[str, Any]]", standard_content_1)
|
||||
)
|
||||
chunk_2 = AIMessageChunk(
|
||||
content=cast("list[str | dict[str, Any]]", standard_content_2)
|
||||
)
|
||||
merged_chunk = chunk_1 + chunk_2
|
||||
assert merged_chunk.content == [
|
||||
{"type": "non_standard", "index": 0, "value": {"foo": "bar baz"}},
|
||||
@@ -467,8 +471,12 @@ def test_content_blocks() -> None:
|
||||
}
|
||||
]
|
||||
standard_content_2 = [{"type": "non_standard", "value": {"foo": "bar"}, "index": 0}]
|
||||
chunk_1 = AIMessageChunk(content=cast("str | list[str | dict]", standard_content_1))
|
||||
chunk_2 = AIMessageChunk(content=cast("str | list[str | dict]", standard_content_2))
|
||||
chunk_1 = AIMessageChunk(
|
||||
content=cast("list[str | dict[str, Any]]", standard_content_1)
|
||||
)
|
||||
chunk_2 = AIMessageChunk(
|
||||
content=cast("list[str | dict[str, Any]]", standard_content_2)
|
||||
)
|
||||
merged_chunk = chunk_1 + chunk_2
|
||||
assert merged_chunk.content == [
|
||||
{
|
||||
|
||||
@@ -758,7 +758,8 @@ class FakeTokenCountingModel(FakeChatModel):
|
||||
def get_num_tokens_from_messages(
|
||||
self,
|
||||
messages: list[BaseMessage],
|
||||
tools: Sequence[dict[str, Any] | type | Callable | BaseTool] | None = None,
|
||||
tools: Sequence[dict[str, Any] | type | Callable[..., Any] | BaseTool]
|
||||
| None = None,
|
||||
) -> int:
|
||||
return dummy_token_counter(messages)
|
||||
|
||||
@@ -1288,7 +1289,7 @@ def test_convert_to_openai_messages_invalid_block() -> None:
|
||||
|
||||
|
||||
def test_handle_openai_responses_blocks() -> None:
|
||||
blocks: str | list[str | dict] = [
|
||||
blocks: str | list[str | dict[str, Any]] = [
|
||||
{"type": "reasoning", "id": "1"},
|
||||
{
|
||||
"type": "function_call",
|
||||
|
||||
@@ -6,7 +6,7 @@ import pytest
|
||||
from pydantic import BaseModel
|
||||
from typing_extensions import override
|
||||
|
||||
from langchain_core.language_models import FakeListChatModel
|
||||
from langchain_core.language_models import FakeListChatModel, LanguageModelInput
|
||||
from langchain_core.load.dump import dumps
|
||||
from langchain_core.load.load import loads
|
||||
from langchain_core.messages import HumanMessage
|
||||
@@ -29,8 +29,8 @@ class FakeStructuredChatModel(FakeListChatModel):
|
||||
|
||||
@override
|
||||
def with_structured_output(
|
||||
self, schema: dict | type[BaseModel], **kwargs: Any
|
||||
) -> Runnable:
|
||||
self, schema: dict[str, Any] | type[BaseModel], **kwargs: Any
|
||||
) -> Runnable[LanguageModelInput, dict[str, Any] | BaseModel]:
|
||||
return RunnableLambda(partial(_fake_runnable, schema=schema, **kwargs))
|
||||
|
||||
@property
|
||||
|
||||
@@ -73,7 +73,7 @@ def _remove_enum(obj: Any) -> None:
|
||||
_remove_enum(item)
|
||||
|
||||
|
||||
def _schema(obj: Any) -> dict:
|
||||
def _schema(obj: Any) -> dict[str, Any]:
|
||||
"""Return the schema of the object."""
|
||||
# Remap to old style schema
|
||||
if isclass(obj):
|
||||
@@ -99,7 +99,7 @@ def _schema(obj: Any) -> dict:
|
||||
raise TypeError(msg)
|
||||
|
||||
|
||||
def _remove_additionalproperties(schema: dict) -> dict[str, Any]:
|
||||
def _remove_additionalproperties(schema: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Remove `"additionalProperties": True` from dicts in the schema.
|
||||
|
||||
Pydantic 2.11 and later versions include `"additionalProperties": True` when
|
||||
|
||||
@@ -335,14 +335,16 @@ class FakeStructuredOutputModel(BaseChatModel):
|
||||
@override
|
||||
def bind_tools(
|
||||
self,
|
||||
tools: Sequence[dict[str, Any] | type[BaseModel] | Callable | BaseTool],
|
||||
tools: Sequence[
|
||||
dict[str, Any] | type[BaseModel] | Callable[..., Any] | BaseTool
|
||||
],
|
||||
**kwargs: Any,
|
||||
) -> Runnable[LanguageModelInput, AIMessage]:
|
||||
return self.bind(tools=tools)
|
||||
|
||||
@override
|
||||
def with_structured_output(
|
||||
self, schema: dict | type[BaseModel], **kwargs: Any
|
||||
self, schema: dict[str, Any] | type[BaseModel], **kwargs: Any
|
||||
) -> Runnable[LanguageModelInput, dict[str, int] | BaseModel]:
|
||||
return RunnableLambda(lambda _: {"foo": self.foo})
|
||||
|
||||
@@ -368,7 +370,9 @@ class FakeModel(BaseChatModel):
|
||||
@override
|
||||
def bind_tools(
|
||||
self,
|
||||
tools: Sequence[dict[str, Any] | type[BaseModel] | Callable | BaseTool],
|
||||
tools: Sequence[
|
||||
dict[str, Any] | type[BaseModel] | Callable[..., Any] | BaseTool
|
||||
],
|
||||
**kwargs: Any,
|
||||
) -> Runnable[LanguageModelInput, AIMessage]:
|
||||
return self.bind(tools=tools)
|
||||
|
||||
@@ -226,7 +226,9 @@ def test_graph_sequence_map(snapshot: SnapshotAssertion) -> None:
|
||||
str_parser = StrOutputParser()
|
||||
xml_parser = XMLOutputParser()
|
||||
|
||||
def conditional_str_parser(value: str) -> Runnable[BaseMessage | str, str]:
|
||||
def conditional_str_parser(
|
||||
value: str,
|
||||
) -> Runnable[BaseMessage | str, str | dict[str, Any]]:
|
||||
if value == "a":
|
||||
return str_parser
|
||||
return xml_parser
|
||||
|
||||
@@ -3770,6 +3770,40 @@ async def test_deep_astream_assign() -> None:
|
||||
}
|
||||
|
||||
|
||||
def _empty_mapper_assign() -> RunnableAssign:
|
||||
"""Build an assign whose mapper yields zero chunks.
|
||||
|
||||
The map output stream is started with `next(map_output, None)` /
|
||||
`anext(map_output, None)`, so `None` is the exhaustion sentinel rather than
|
||||
a real chunk. Both `stream` and `astream` must guard against yielding that
|
||||
sentinel into the output stream.
|
||||
"""
|
||||
|
||||
def empty_gen(it: Iterator[Any]) -> Iterator[dict[str, Any]]:
|
||||
for _ in it:
|
||||
pass
|
||||
yield from () # consume input, yield nothing
|
||||
|
||||
async def aempty_gen(it: AsyncIterator[Any]) -> AsyncIterator[dict[str, Any]]:
|
||||
async for _ in it:
|
||||
pass
|
||||
return
|
||||
yield # pragma: no cover # make this an async generator function
|
||||
|
||||
return RunnablePassthrough.assign(foo=RunnableGenerator(empty_gen, aempty_gen))
|
||||
|
||||
|
||||
def test_stream_assign_empty_mapper() -> None:
|
||||
"""An assign whose mapper yields no chunks must not emit `None` (sync)."""
|
||||
assert list(_empty_mapper_assign().stream({"a": 1})) == [{"a": 1}]
|
||||
|
||||
|
||||
async def test_astream_assign_empty_mapper() -> None:
|
||||
"""An assign whose mapper yields no chunks must not emit `None` (async)."""
|
||||
chunks = [chunk async for chunk in _empty_mapper_assign().astream({"a": 1})]
|
||||
assert chunks == [{"a": 1}]
|
||||
|
||||
|
||||
def test_runnable_sequence_transform() -> None:
|
||||
llm = FakeStreamingListLLM(responses=["foo-lish"])
|
||||
|
||||
@@ -5374,7 +5408,7 @@ async def test_ainvoke_on_returned_runnable() -> None:
|
||||
|
||||
|
||||
def test_invoke_stream_passthrough_assign_trace() -> None:
|
||||
def idchain_sync(_input: dict, /) -> bool:
|
||||
def idchain_sync(_input: dict[str, Any], /) -> bool:
|
||||
return False
|
||||
|
||||
chain = RunnablePassthrough.assign(urls=idchain_sync)
|
||||
@@ -5394,7 +5428,7 @@ def test_invoke_stream_passthrough_assign_trace() -> None:
|
||||
|
||||
|
||||
async def test_ainvoke_astream_passthrough_assign_trace() -> None:
|
||||
def idchain_sync(_input: dict, /) -> bool:
|
||||
def idchain_sync(_input: dict[str, Any], /) -> bool:
|
||||
return False
|
||||
|
||||
chain = RunnablePassthrough.assign(urls=idchain_sync)
|
||||
|
||||
@@ -78,7 +78,7 @@ def _assert_events_equal_allow_superset_metadata(
|
||||
async def test_event_stream_with_simple_function_tool() -> None:
|
||||
"""Test the event stream with a function and tool."""
|
||||
|
||||
def foo(_: int) -> dict:
|
||||
def foo(_: int) -> dict[str, Any]:
|
||||
"""Foo."""
|
||||
return {"x": 5}
|
||||
|
||||
@@ -1076,12 +1076,14 @@ async def test_event_streaming_with_tools() -> None:
|
||||
return "world"
|
||||
|
||||
@tool
|
||||
def with_parameters(x: int, y: str) -> dict:
|
||||
def with_parameters(x: int, y: str) -> dict[str, Any]:
|
||||
"""A tool that does nothing."""
|
||||
return {"x": x, "y": y}
|
||||
|
||||
@tool
|
||||
def with_parameters_and_callbacks(x: int, y: str, callbacks: Callbacks) -> dict:
|
||||
def with_parameters_and_callbacks(
|
||||
x: int, y: str, callbacks: Callbacks
|
||||
) -> dict[str, Any]:
|
||||
"""A tool that does nothing."""
|
||||
_ = callbacks
|
||||
return {"x": x, "y": y}
|
||||
|
||||
@@ -93,7 +93,7 @@ async def _collect_events(
|
||||
async def test_event_stream_with_simple_function_tool() -> None:
|
||||
"""Test the event stream with a function and tool."""
|
||||
|
||||
def foo(x: int) -> dict:
|
||||
def foo(x: int) -> dict[str, Any]:
|
||||
"""Foo."""
|
||||
_ = x
|
||||
return {"x": 5}
|
||||
@@ -1097,12 +1097,14 @@ async def test_event_streaming_with_tools() -> None:
|
||||
return "world"
|
||||
|
||||
@tool
|
||||
def with_parameters(x: int, y: str) -> dict:
|
||||
def with_parameters(x: int, y: str) -> dict[str, Any]:
|
||||
"""A tool that does nothing."""
|
||||
return {"x": x, "y": y}
|
||||
|
||||
@tool
|
||||
def with_parameters_and_callbacks(x: int, y: str, callbacks: Callbacks) -> dict:
|
||||
def with_parameters_and_callbacks(
|
||||
x: int, y: str, callbacks: Callbacks
|
||||
) -> dict[str, Any]:
|
||||
"""A tool that does nothing."""
|
||||
_ = callbacks
|
||||
return {"x": x, "y": y}
|
||||
|
||||
@@ -421,7 +421,8 @@ class TestRunnableSequenceParallelTraceNesting:
|
||||
ids=["invoke", "stream", "batch"],
|
||||
)
|
||||
def test_sync(
|
||||
self, method: Callable[[RunnableLambda, list[BaseCallbackHandler]], int]
|
||||
self,
|
||||
method: Callable[[RunnableLambda[int, int], list[BaseCallbackHandler]], int],
|
||||
) -> None:
|
||||
def other_thing(_: int) -> Generator[int, None, None]:
|
||||
yield 1
|
||||
@@ -458,7 +459,8 @@ class TestRunnableSequenceParallelTraceNesting:
|
||||
async def test_async(
|
||||
self,
|
||||
method: Callable[
|
||||
[RunnableLambda, list[BaseCallbackHandler]], Coroutine[Any, Any, int]
|
||||
[RunnableLambda[int, int], list[BaseCallbackHandler]],
|
||||
Coroutine[Any, Any, int],
|
||||
],
|
||||
) -> None:
|
||||
async def other_thing(_: int) -> AsyncGenerator[int, None]:
|
||||
|
||||
@@ -56,7 +56,7 @@ def test_nonlocals() -> None:
|
||||
def my_func4(value: str) -> str:
|
||||
return global_agent.invoke(value)
|
||||
|
||||
def my_func5() -> tuple[Callable[[str], str], RunnableLambda]:
|
||||
def my_func5() -> tuple[Callable[[str], str], RunnableLambda[str, str]]:
|
||||
global_agent = RunnableLambda[str, str](lambda x: x * 3)
|
||||
|
||||
def my_func6(value: str) -> str:
|
||||
|
||||
@@ -23,7 +23,7 @@ class TestSyncInMemoryStore(BaseStoreSyncTests[Any]):
|
||||
return "value1", "value2", "value3"
|
||||
|
||||
|
||||
class TestAsyncInMemoryStore(BaseStoreAsyncTests):
|
||||
class TestAsyncInMemoryStore(BaseStoreAsyncTests[Any]):
|
||||
@pytest.fixture
|
||||
@override
|
||||
async def kv_store(self) -> InMemoryStore:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import uuid
|
||||
from typing import get_args
|
||||
from typing import Any, get_args
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -488,7 +488,7 @@ def test_message_chunk_to_message() -> None:
|
||||
|
||||
|
||||
def test_tool_calls_merge() -> None:
|
||||
chunks: list[dict] = [
|
||||
chunks: list[dict[str, Any]] = [
|
||||
{"content": ""},
|
||||
{
|
||||
"content": "",
|
||||
@@ -1092,7 +1092,11 @@ def test_tool_message_str() -> None:
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_merge_content(first: list | str, others: list, expected: list | str) -> None:
|
||||
def test_merge_content(
|
||||
first: str | list[str | dict[str, Any]],
|
||||
others: str | list[str | dict[str, Any]],
|
||||
expected: str | list[str | dict[str, Any]],
|
||||
) -> None:
|
||||
actual = merge_content(first, *others)
|
||||
assert actual == expected
|
||||
|
||||
|
||||
@@ -12,7 +12,7 @@ from langchain_core.prompt_values import ChatPromptValueConcrete
|
||||
|
||||
|
||||
def test_chat_prompt_value_concrete() -> None:
|
||||
messages: list = [
|
||||
messages = [
|
||||
AIMessage("foo"),
|
||||
HumanMessage("foo"),
|
||||
SystemMessage("foo"),
|
||||
|
||||
@@ -113,7 +113,7 @@ class _MockSchema(BaseModel):
|
||||
|
||||
arg1: int
|
||||
arg2: bool
|
||||
arg3: dict | None = None
|
||||
arg3: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class _MockStructuredTool(BaseTool):
|
||||
@@ -122,10 +122,12 @@ class _MockStructuredTool(BaseTool):
|
||||
description: str = "A Structured Tool"
|
||||
|
||||
@override
|
||||
def _run(self, *, arg1: int, arg2: bool, arg3: dict | None = None) -> str:
|
||||
def _run(self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None) -> str:
|
||||
return f"{arg1} {arg2} {arg3}"
|
||||
|
||||
async def _arun(self, *, arg1: int, arg2: bool, arg3: dict | None = None) -> str:
|
||||
async def _arun(
|
||||
self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
|
||||
) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@@ -166,11 +168,13 @@ def test_misannotated_base_tool_raises_error() -> None:
|
||||
description: str = "A Structured Tool"
|
||||
|
||||
@override
|
||||
def _run(self, *, arg1: int, arg2: bool, arg3: dict | None = None) -> str:
|
||||
def _run(
|
||||
self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
|
||||
) -> str:
|
||||
return f"{arg1} {arg2} {arg3}"
|
||||
|
||||
async def _arun(
|
||||
self, *, arg1: int, arg2: bool, arg3: dict | None = None
|
||||
self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
|
||||
) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -184,11 +188,13 @@ def test_forward_ref_annotated_base_tool_accepted() -> None:
|
||||
description: str = "A Structured Tool"
|
||||
|
||||
@override
|
||||
def _run(self, *, arg1: int, arg2: bool, arg3: dict | None = None) -> str:
|
||||
def _run(
|
||||
self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
|
||||
) -> str:
|
||||
return f"{arg1} {arg2} {arg3}"
|
||||
|
||||
async def _arun(
|
||||
self, *, arg1: int, arg2: bool, arg3: dict | None = None
|
||||
self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
|
||||
) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -202,11 +208,13 @@ def test_subclass_annotated_base_tool_accepted() -> None:
|
||||
description: str = "A Structured Tool"
|
||||
|
||||
@override
|
||||
def _run(self, *, arg1: int, arg2: bool, arg3: dict | None = None) -> str:
|
||||
def _run(
|
||||
self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
|
||||
) -> str:
|
||||
return f"{arg1} {arg2} {arg3}"
|
||||
|
||||
async def _arun(
|
||||
self, *, arg1: int, arg2: bool, arg3: dict | None = None
|
||||
self, *, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
|
||||
) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -219,7 +227,7 @@ def test_decorator_with_specified_schema() -> None:
|
||||
"""Test that manually specified schemata are passed through to the tool."""
|
||||
|
||||
@tool(args_schema=_MockSchema)
|
||||
def tool_func(*, arg1: int, arg2: bool, arg3: dict | None = None) -> str:
|
||||
def tool_func(*, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None) -> str:
|
||||
return f"{arg1} {arg2} {arg3}"
|
||||
|
||||
assert isinstance(tool_func, BaseTool)
|
||||
@@ -238,10 +246,12 @@ def test_decorator_with_specified_schema_pydantic_v1() -> None:
|
||||
|
||||
arg1: int
|
||||
arg2: bool
|
||||
arg3: dict | None = None
|
||||
arg3: dict[str, Any] | None = None
|
||||
|
||||
@tool(args_schema=cast("ArgsSchema", _MockSchemaV1))
|
||||
def tool_func_v1(*, arg1: int, arg2: bool, arg3: dict | None = None) -> str:
|
||||
def tool_func_v1(
|
||||
*, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
|
||||
) -> str:
|
||||
return f"{arg1} {arg2} {arg3}"
|
||||
|
||||
assert isinstance(tool_func_v1, BaseTool)
|
||||
@@ -253,7 +263,7 @@ def test_decorated_function_schema_equivalent() -> None:
|
||||
|
||||
@tool
|
||||
def structured_tool_input(
|
||||
*, arg1: int, arg2: bool, arg3: dict | None = None
|
||||
*, arg1: int, arg2: bool, arg3: dict[str, Any] | None = None
|
||||
) -> str:
|
||||
"""Return the arguments directly."""
|
||||
return f"{arg1} {arg2} {arg3}"
|
||||
@@ -322,7 +332,7 @@ def test_structured_args_decorator_no_infer_schema() -> None:
|
||||
|
||||
@tool(infer_schema=False)
|
||||
def structured_tool_input(
|
||||
arg1: int, arg2: float | datetime, opt_arg: dict | None = None
|
||||
arg1: int, arg2: float | datetime, opt_arg: dict[str, Any] | None = None
|
||||
) -> str:
|
||||
"""Return the arguments directly."""
|
||||
return f"{arg1}, {arg2}, {opt_arg}"
|
||||
@@ -362,7 +372,7 @@ def test_structured_tool_types_parsed() -> None:
|
||||
def structured_tool(
|
||||
some_enum: SomeEnum,
|
||||
some_base_model: SomeBaseModel,
|
||||
) -> dict:
|
||||
) -> dict[str, Any]:
|
||||
"""Return the arguments directly."""
|
||||
return {
|
||||
"some_enum": some_enum,
|
||||
@@ -1032,7 +1042,7 @@ def test_validation_error_handling_non_validation_error(
|
||||
|
||||
def _parse_input(
|
||||
self,
|
||||
tool_input: str | dict,
|
||||
tool_input: str | dict[str, Any],
|
||||
tool_call_id: str | None,
|
||||
) -> str | dict[str, Any]:
|
||||
raise NotImplementedError
|
||||
@@ -1098,7 +1108,7 @@ async def test_async_validation_error_handling_non_validation_error(
|
||||
|
||||
def _parse_input(
|
||||
self,
|
||||
tool_input: str | dict,
|
||||
tool_input: str | dict[str, Any],
|
||||
tool_call_id: str | None,
|
||||
) -> str | dict[str, Any]:
|
||||
raise NotImplementedError
|
||||
@@ -1145,9 +1155,11 @@ def test_optional_subset_model_rewrite() -> None:
|
||||
({"bar": "bar", "baz": None}, {"bar": "bar", "baz": None, "buzz": "buzz"}),
|
||||
],
|
||||
)
|
||||
def test_tool_invoke_optional_args(inputs: dict, expected: dict | None) -> None:
|
||||
def test_tool_invoke_optional_args(
|
||||
inputs: dict[str, Any], expected: dict[str, Any] | None
|
||||
) -> None:
|
||||
@tool
|
||||
def foo(bar: str, baz: int | None = 3, buzz: str | None = "buzz") -> dict:
|
||||
def foo(bar: str, baz: int | None = 3, buzz: str | None = "buzz") -> dict[str, Any]:
|
||||
"""The foo."""
|
||||
return {
|
||||
"bar": bar,
|
||||
@@ -2283,9 +2295,13 @@ def test__get_all_basemodel_annotations_v2(*, use_v1_namespace: bool) -> None:
|
||||
return "foo"
|
||||
|
||||
class ModelC(Mixin, ModelB):
|
||||
c: dict
|
||||
c: dict[str, Any]
|
||||
|
||||
expected = {"a": str, "b": Annotated[ModelA[dict[str, Any]], "foo"], "c": dict}
|
||||
expected = {
|
||||
"a": str,
|
||||
"b": Annotated[ModelA[dict[str, Any]], "foo"],
|
||||
"c": dict[str, Any],
|
||||
}
|
||||
actual = get_all_basemodel_annotations(ModelC)
|
||||
assert actual == expected
|
||||
|
||||
@@ -2309,7 +2325,7 @@ def test__get_all_basemodel_annotations_v2(*, use_v1_namespace: bool) -> None:
|
||||
expected = {
|
||||
"a": str,
|
||||
"b": Annotated[ModelA[dict[str, Any]], "foo"],
|
||||
"c": dict,
|
||||
"c": dict[str, Any],
|
||||
"d": str | int | None,
|
||||
}
|
||||
actual = get_all_basemodel_annotations(ModelD)
|
||||
@@ -2318,7 +2334,7 @@ def test__get_all_basemodel_annotations_v2(*, use_v1_namespace: bool) -> None:
|
||||
expected = {
|
||||
"a": str,
|
||||
"b": Annotated[ModelA[dict[str, Any]], "foo"],
|
||||
"c": dict,
|
||||
"c": dict[str, Any],
|
||||
"d": int | None,
|
||||
}
|
||||
actual = get_all_basemodel_annotations(ModelD[int])
|
||||
@@ -2342,7 +2358,7 @@ def test_tool_annotations_preserved() -> None:
|
||||
"""Test that annotations are preserved when creating a tool."""
|
||||
|
||||
@tool
|
||||
def my_tool(val: int, other_val: Annotated[dict, "my annotation"]) -> str:
|
||||
def my_tool(val: int, other_val: Annotated[dict[str, Any], "my annotation"]) -> str:
|
||||
"""Tool docstring."""
|
||||
return "foo"
|
||||
|
||||
@@ -3080,7 +3096,7 @@ def test_tool_args_schema_with_annotated_type() -> None:
|
||||
class CallbackHandlerWithInputCapture(FakeCallbackHandler):
|
||||
"""Callback handler that captures inputs passed to on_tool_start."""
|
||||
|
||||
captured_inputs: list[dict | None] = Field(default_factory=list)
|
||||
captured_inputs: list[dict[str, Any] | None] = Field(default_factory=list)
|
||||
|
||||
def on_tool_start(
|
||||
self,
|
||||
@@ -3114,7 +3130,7 @@ def test_filter_injected_args_from_callbacks() -> None:
|
||||
@tool
|
||||
def search_tool(
|
||||
query: str,
|
||||
state: Annotated[dict, InjectedToolArg()],
|
||||
state: Annotated[dict[str, Any], InjectedToolArg()],
|
||||
) -> str:
|
||||
"""Search with injected state.
|
||||
|
||||
@@ -3182,7 +3198,7 @@ def test_filter_multiple_injected_args() -> None:
|
||||
def complex_tool(
|
||||
query: str,
|
||||
limit: int,
|
||||
state: Annotated[dict, InjectedToolArg()],
|
||||
state: Annotated[dict[str, Any], InjectedToolArg()],
|
||||
context: Annotated[str, InjectedToolArg()],
|
||||
run_manager: CallbackManagerForToolRun | None = None,
|
||||
) -> str:
|
||||
@@ -3250,7 +3266,7 @@ async def test_filter_injected_args_async() -> None:
|
||||
@tool
|
||||
async def async_search_tool(
|
||||
query: str,
|
||||
state: Annotated[dict, InjectedToolArg()],
|
||||
state: Annotated[dict[str, Any], InjectedToolArg()],
|
||||
) -> str:
|
||||
"""Async search with injected state.
|
||||
|
||||
|
||||
@@ -12,7 +12,7 @@ async def test_same_event_loop() -> None:
|
||||
This is the easy case.
|
||||
"""
|
||||
reader_loop = asyncio.get_event_loop()
|
||||
channel = _MemoryStream[dict](reader_loop)
|
||||
channel = _MemoryStream[dict[str, int | float]](reader_loop)
|
||||
writer = channel.get_send_stream()
|
||||
reader = channel.get_receive_stream()
|
||||
|
||||
@@ -30,7 +30,7 @@ async def test_same_event_loop() -> None:
|
||||
)
|
||||
await writer.aclose()
|
||||
|
||||
async def consumer() -> AsyncIterator[dict]:
|
||||
async def consumer() -> AsyncIterator[dict[str, int | float]]:
|
||||
tic = time.time()
|
||||
async for item in reader:
|
||||
toc = time.time()
|
||||
@@ -63,7 +63,7 @@ async def test_same_event_loop() -> None:
|
||||
async def test_queue_for_streaming_via_sync_call() -> None:
|
||||
"""Test via async -> sync -> async path."""
|
||||
reader_loop = asyncio.get_event_loop()
|
||||
channel = _MemoryStream[dict](reader_loop)
|
||||
channel = _MemoryStream[dict[str, int | float]](reader_loop)
|
||||
writer = channel.get_send_stream()
|
||||
reader = channel.get_receive_stream()
|
||||
|
||||
@@ -85,7 +85,7 @@ async def test_queue_for_streaming_via_sync_call() -> None:
|
||||
"""Blocking sync call."""
|
||||
asyncio.run(producer())
|
||||
|
||||
async def consumer() -> AsyncIterator[dict]:
|
||||
async def consumer() -> AsyncIterator[dict[str, int | float]]:
|
||||
tic = time.time()
|
||||
async for item in reader:
|
||||
toc = time.time()
|
||||
|
||||
@@ -26,7 +26,7 @@ from pydantic import BaseModel, ConfigDict, Field
|
||||
from pydantic.errors import PydanticInvalidForJsonSchema
|
||||
|
||||
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
||||
from langchain_core.runnables import Runnable, RunnableLambda
|
||||
from langchain_core.runnables import RunnableLambda
|
||||
from langchain_core.tools import BaseTool, StructuredTool, Tool, tool
|
||||
from langchain_core.utils.function_calling import (
|
||||
_convert_typed_dict_to_openai_function,
|
||||
@@ -49,7 +49,7 @@ def pydantic() -> type[BaseModel]:
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def annotated_function() -> Callable:
|
||||
def annotated_function() -> Callable[[int, Literal["bar", "baz"]], None]:
|
||||
def dummy_function(
|
||||
arg1: ExtensionsAnnotated[int, "foo"],
|
||||
arg2: ExtensionsAnnotated[Literal["bar", "baz"], "one of 'bar', 'baz'"],
|
||||
@@ -60,7 +60,7 @@ def annotated_function() -> Callable:
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def function() -> Callable:
|
||||
def function() -> Callable[[int, Literal["bar", "baz"]], None]:
|
||||
def dummy_function(arg1: int, arg2: Literal["bar", "baz"]) -> None:
|
||||
"""Dummy function.
|
||||
|
||||
@@ -73,7 +73,7 @@ def function() -> Callable:
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def function_docstring_annotations() -> Callable:
|
||||
def function_docstring_annotations() -> Callable[[int, Literal["bar", "baz"]], None]:
|
||||
def dummy_function(arg1: int, arg2: Literal["bar", "baz"]) -> None:
|
||||
"""Dummy function.
|
||||
|
||||
@@ -85,13 +85,14 @@ def function_docstring_annotations() -> Callable:
|
||||
return dummy_function
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def runnable() -> Runnable:
|
||||
class Args(ExtensionsTypedDict):
|
||||
arg1: ExtensionsAnnotated[int, "foo"]
|
||||
arg2: ExtensionsAnnotated[Literal["bar", "baz"], "one of 'bar', 'baz'"]
|
||||
class _Args(ExtensionsTypedDict):
|
||||
arg1: ExtensionsAnnotated[int, "foo"]
|
||||
arg2: ExtensionsAnnotated[Literal["bar", "baz"], "one of 'bar', 'baz'"]
|
||||
|
||||
def dummy_function(input_dict: Args) -> None:
|
||||
|
||||
@pytest.fixture
|
||||
def runnable() -> RunnableLambda[_Args, None]:
|
||||
def dummy_function(input_dict: _Args) -> None:
|
||||
pass
|
||||
|
||||
return RunnableLambda(dummy_function)
|
||||
@@ -229,7 +230,7 @@ def dummy_extensions_typed_dict_docstring() -> type:
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def json_schema() -> dict:
|
||||
def json_schema() -> dict[str, Any]:
|
||||
return {
|
||||
"title": "dummy_function",
|
||||
"description": "Dummy function.",
|
||||
@@ -247,7 +248,7 @@ def json_schema() -> dict:
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def anthropic_tool() -> dict:
|
||||
def anthropic_tool() -> dict[str, Any]:
|
||||
return {
|
||||
"name": "dummy_function",
|
||||
"description": "Dummy function.",
|
||||
@@ -267,7 +268,7 @@ def anthropic_tool() -> dict:
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def bedrock_converse_tool() -> dict:
|
||||
def bedrock_converse_tool() -> dict[str, Any]:
|
||||
return {
|
||||
"toolSpec": {
|
||||
"name": "dummy_function",
|
||||
@@ -313,17 +314,17 @@ class DummyWithClassMethod:
|
||||
|
||||
def test_convert_to_openai_function(
|
||||
pydantic: type[BaseModel],
|
||||
function: Callable,
|
||||
function_docstring_annotations: Callable,
|
||||
function: Callable[[int, Literal["bar", "baz"]], None],
|
||||
function_docstring_annotations: Callable[[int, Literal["bar", "baz"]], None],
|
||||
dummy_structured_tool: StructuredTool,
|
||||
dummy_structured_tool_args_schema_dict: StructuredTool,
|
||||
dummy_tool: BaseTool,
|
||||
json_schema: dict,
|
||||
anthropic_tool: dict,
|
||||
bedrock_converse_tool: dict,
|
||||
annotated_function: Callable,
|
||||
json_schema: dict[str, Any],
|
||||
anthropic_tool: dict[str, Any],
|
||||
bedrock_converse_tool: dict[str, Any],
|
||||
annotated_function: Callable[[int, Literal["bar", "baz"]], None],
|
||||
dummy_pydantic: type[BaseModel],
|
||||
runnable: Runnable,
|
||||
runnable: RunnableLambda[_Args, None],
|
||||
dummy_typing_typed_dict: type,
|
||||
dummy_typing_typed_dict_docstring: type,
|
||||
dummy_extensions_typed_dict: type,
|
||||
@@ -643,7 +644,7 @@ openai_function_no_description_no_params = {
|
||||
openai_function_no_description,
|
||||
],
|
||||
)
|
||||
def test_convert_to_openai_function_no_description(func: dict) -> None:
|
||||
def test_convert_to_openai_function_no_description(func: dict[str, Any]) -> None:
|
||||
expected = {
|
||||
"name": "dummy_function",
|
||||
"parameters": {
|
||||
@@ -670,7 +671,9 @@ def test_convert_to_openai_function_no_description(func: dict) -> None:
|
||||
openai_function_no_description_no_params,
|
||||
],
|
||||
)
|
||||
def test_convert_to_openai_function_no_description_no_params(func: dict) -> None:
|
||||
def test_convert_to_openai_function_no_description_no_params(
|
||||
func: dict[str, Any],
|
||||
) -> None:
|
||||
expected = {
|
||||
"name": "dummy_function",
|
||||
}
|
||||
@@ -1040,7 +1043,9 @@ def test__convert_typed_dict_to_openai_function(
|
||||
@pytest.mark.parametrize("typed_dict", [ExtensionsTypedDict, TypingTypedDict])
|
||||
def test__convert_typed_dict_to_openai_function_fail(typed_dict: type) -> None:
|
||||
class Tool(typed_dict): # type: ignore[misc]
|
||||
arg1: typing.MutableSet # Pydantic 2 supports this, but pydantic v1 does not.
|
||||
arg1: typing.MutableSet[
|
||||
Any
|
||||
] # Pydantic 2 supports this, but pydantic v1 does not.
|
||||
|
||||
# Error should be raised since we're using v1 code path here
|
||||
with pytest.raises(TypeError):
|
||||
@@ -1081,15 +1086,15 @@ def test_convert_to_openai_function_no_args() -> None:
|
||||
|
||||
def test_convert_to_json_schema(
|
||||
pydantic: type[BaseModel],
|
||||
function: Callable,
|
||||
function_docstring_annotations: Callable,
|
||||
function: Callable[[int, Literal["bar", "baz"]], None],
|
||||
function_docstring_annotations: Callable[[int, Literal["bar", "baz"]], None],
|
||||
dummy_structured_tool: StructuredTool,
|
||||
dummy_structured_tool_args_schema_dict: StructuredTool,
|
||||
dummy_tool: BaseTool,
|
||||
json_schema: dict,
|
||||
anthropic_tool: dict,
|
||||
bedrock_converse_tool: dict,
|
||||
annotated_function: Callable,
|
||||
json_schema: dict[str, Any],
|
||||
anthropic_tool: dict[str, Any],
|
||||
bedrock_converse_tool: dict[str, Any],
|
||||
annotated_function: Callable[[int, Literal["bar", "baz"]], None],
|
||||
dummy_pydantic: type[BaseModel],
|
||||
dummy_typing_typed_dict: type,
|
||||
dummy_typing_typed_dict_docstring: type,
|
||||
@@ -1123,10 +1128,10 @@ def test_convert_to_json_schema(
|
||||
|
||||
|
||||
def test_convert_to_openai_function_nested_strict_2() -> None:
|
||||
def my_function(arg1: dict, arg2: dict | None) -> None:
|
||||
def my_function(arg1: dict[str, Any], arg2: dict[str, Any] | None) -> None:
|
||||
"""Dummy function."""
|
||||
|
||||
expected: dict = {
|
||||
expected: dict[str, Any] = {
|
||||
"name": "my_function",
|
||||
"description": "Dummy function.",
|
||||
"parameters": {
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from langchain_core.utils.function_calling import _rm_titles
|
||||
@@ -229,5 +231,5 @@ output5 = {
|
||||
(schema5, output5),
|
||||
],
|
||||
)
|
||||
def test_rm_titles(schema: dict, output: dict) -> None:
|
||||
def test_rm_titles(schema: dict[str, Any], output: dict[str, Any]) -> None:
|
||||
assert _rm_titles(schema) == output
|
||||
@@ -132,7 +132,9 @@ def test_check_package_version(
|
||||
],
|
||||
)
|
||||
def test_merge_dicts(
|
||||
left: dict, right: dict, expected: dict | AbstractContextManager
|
||||
left: dict[str, Any],
|
||||
right: dict[str, Any],
|
||||
expected: dict[str, Any] | AbstractContextManager[BaseException],
|
||||
) -> None:
|
||||
err = expected if isinstance(expected, AbstractContextManager) else nullcontext()
|
||||
|
||||
@@ -160,7 +162,9 @@ def test_merge_dicts(
|
||||
)
|
||||
@pytest.mark.xfail(reason="Refactors to make in 0.3")
|
||||
def test_merge_dicts_0_3(
|
||||
left: dict, right: dict, expected: dict | AbstractContextManager
|
||||
left: dict[str, Any],
|
||||
right: dict[str, Any],
|
||||
expected: dict[str, Any] | AbstractContextManager[BaseException],
|
||||
) -> None:
|
||||
err = expected if isinstance(expected, AbstractContextManager) else nullcontext()
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ import pytest
|
||||
|
||||
pytest.importorskip("numpy")
|
||||
import numpy as np
|
||||
import numpy.typing as npt
|
||||
|
||||
from langchain_core.vectorstores.utils import _cosine_similarity
|
||||
|
||||
@@ -61,8 +62,8 @@ class TestCosineSimilarity:
|
||||
|
||||
def test_numpy_array_input(self) -> None:
|
||||
"""Test with numpy array inputs."""
|
||||
x: np.ndarray = np.array([[1, 0], [0, 1]])
|
||||
y: np.ndarray = np.array([[1, 0], [0, 1]])
|
||||
x: npt.NDArray[np.floating] = np.array([[1, 0], [0, 1]])
|
||||
y: npt.NDArray[np.floating] = np.array([[1, 0], [0, 1]])
|
||||
result = _cosine_similarity(x, y)
|
||||
expected = np.array([[1.0, 0.0], [0.0, 1.0]])
|
||||
np.testing.assert_array_almost_equal(result, expected)
|
||||
@@ -70,7 +71,7 @@ class TestCosineSimilarity:
|
||||
def test_mixed_input_types(self) -> None:
|
||||
"""Test with mixed input types (list and numpy array)."""
|
||||
x: list[list[float]] = [[1, 0], [0, 1]]
|
||||
y: np.ndarray = np.array([[1, 0], [0, 1]])
|
||||
y: npt.NDArray[np.floating] = np.array([[1, 0], [0, 1]])
|
||||
result = _cosine_similarity(x, y)
|
||||
expected = np.array([[1.0, 0.0], [0.0, 1.0]])
|
||||
np.testing.assert_array_almost_equal(result, expected)
|
||||
|
||||
@@ -30,7 +30,7 @@ class CustomAddTextsVectorstore(VectorStore):
|
||||
def add_texts(
|
||||
self,
|
||||
texts: Iterable[str],
|
||||
metadatas: list[dict] | None = None,
|
||||
metadatas: list[dict[str, Any]] | None = None,
|
||||
ids: list[str] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> list[str]:
|
||||
@@ -58,7 +58,7 @@ class CustomAddTextsVectorstore(VectorStore):
|
||||
cls,
|
||||
texts: list[str],
|
||||
embedding: Embeddings,
|
||||
metadatas: list[dict] | None = None,
|
||||
metadatas: list[dict[str, Any]] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> CustomAddTextsVectorstore:
|
||||
vectorstore = CustomAddTextsVectorstore()
|
||||
@@ -104,7 +104,7 @@ class CustomAddDocumentsVectorstore(VectorStore):
|
||||
cls,
|
||||
texts: list[str],
|
||||
embedding: Embeddings,
|
||||
metadatas: list[dict] | None = None,
|
||||
metadatas: list[dict[str, Any]] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> CustomAddDocumentsVectorstore:
|
||||
vectorstore = CustomAddDocumentsVectorstore()
|
||||
|
||||
@@ -455,7 +455,8 @@ class ExperimentalMarkdownSyntaxTextSplitter:
|
||||
# Apply the header stack as metadata
|
||||
for depth, value in self.current_header_stack:
|
||||
header_key = self.splittable_headers.get("#" * depth)
|
||||
self.current_chunk.metadata[header_key] = value
|
||||
if header_key is not None:
|
||||
self.current_chunk.metadata[header_key] = value
|
||||
self.chunks.append(self.current_chunk)
|
||||
# Reset the current chunk
|
||||
self.current_chunk = Document(page_content="")
|
||||
|
||||
Reference in new issue
Block a user