Support index embed specification via string (#3317)

So you can do

```
from langgraph.store.memory|posgres|etc. import InMemoryStore

InMemoryStore(index={"embed": "openai:text-embedding-3-small"})
```
This commit is contained in:
William FH
2025-02-09 02:09:10 +00:00
committed by GitHub
parent 16cfeff78c
commit e81979827f
4 changed files with 44 additions and 4 deletions
@@ -493,13 +493,14 @@ class IndexConfig(TypedDict, total=False):
- cohere:embed-multilingual-light-v3.0: 384
"""
embed: Union[Embeddings, EmbeddingsFunc, AEmbeddingsFunc]
embed: Union[Embeddings, EmbeddingsFunc, AEmbeddingsFunc, str]
"""Optional function to generate embeddings from text.
Can be specified in three ways:
1. A LangChain Embeddings instance
2. A synchronous embedding function (EmbeddingsFunc)
3. An asynchronous embedding function (AEmbeddingsFunc)
4. A provider string (e.g., "openai:text-embedding-3-small")
???+ example "Examples"
Using LangChain's initialization with InMemoryStore:
+40 -1
View File
@@ -7,6 +7,7 @@ asynchronous operations.
"""
import asyncio
import functools
import json
from typing import Any, Awaitable, Callable, Optional, Sequence, Union
@@ -28,7 +29,7 @@ Similar to EmbeddingsFunc, but returns an awaitable that resolves to the embeddi
def ensure_embeddings(
embed: Union[Embeddings, EmbeddingsFunc, AEmbeddingsFunc, None],
embed: Union[Embeddings, EmbeddingsFunc, AEmbeddingsFunc, str, None],
) -> Embeddings:
"""Ensure that an embedding function conforms to LangChain's Embeddings interface.
@@ -62,9 +63,37 @@ def ensure_embeddings(
embeddings = ensure_embeddings(my_async_fn)
result = await embeddings.aembed_query("hello") # Returns [0.1, 0.2]
```
Initialize embeddings using a provider string:
```python
# Requires langchain>=0.3.9 and langgraph-checkpoint>=2.0.11
embeddings = ensure_embeddings("openai:text-embedding-3-small")
result = embeddings.embed_query("hello")
```
"""
if embed is None:
raise ValueError("embed must be provided")
if isinstance(embed, str):
init_embeddings = _get_init_embeddings()
if init_embeddings is None:
from importlib.metadata import PackageNotFoundError, version
try:
lc_version = version("langchain")
version_info = f"Found langchain version {lc_version}, but"
except PackageNotFoundError:
version_info = "langchain is not installed;"
raise ValueError(
f"Could not load embeddings from string '{embed}'. {version_info} "
"loading embeddings by provider:identifier string requires langchain>=0.3.9 "
"as well as the provider-specific package. "
"Install LangChain with: pip install 'langchain>=0.3.9' "
"and the provider-specific package (e.g., 'langchain-openai>=0.3.0'). "
"Alternatively, specify 'embed' as a compatible Embeddings object or python function."
)
return init_embeddings(embed)
if isinstance(embed, Embeddings):
return embed
return EmbeddingsLambda(embed)
@@ -373,6 +402,16 @@ def _is_async_callable(
)
@functools.lru_cache
def _get_init_embeddings() -> Optional[Callable[[str], Embeddings]]:
try:
from langchain.embeddings import init_embeddings # type: ignore
return init_embeddings
except ImportError:
return None
__all__ = [
"ensure_embeddings",
"EmbeddingsFunc",
@@ -493,7 +493,7 @@ def _cosine_similarity(X: list[float], Y: list[list[float]]) -> list[float]:
if not Y:
return []
if _check_numpy():
import numpy as np # type: ignore
import numpy as np # type: ignore[import-not-found]
X_arr = np.array(X) if not isinstance(X, np.ndarray) else X
Y_arr = np.array(Y) if not isinstance(Y, np.ndarray) else Y
+1 -1
View File
@@ -1,6 +1,6 @@
[tool.poetry]
name = "langgraph-checkpoint"
version = "2.0.11"
version = "2.0.12"
description = "Library with base interfaces for LangGraph checkpoint savers."
authors = []
license = "MIT"