Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 31 additions & 17 deletions redisvl/utils/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
from enum import Enum
from functools import wraps
from time import time
from typing import Any, Callable, Coroutine, Sequence, TypeVar
from typing import Any, Callable, Coroutine, ParamSpec, Sequence, TypeVar, cast
from warnings import warn

from pydantic import BaseModel
Expand All @@ -17,6 +17,9 @@
from redisvl.types import SyncRedisClient

T = TypeVar("T")
R = TypeVar("R")
P = ParamSpec("P")
C = TypeVar("C", bound=type)


def create_ulid() -> str:
Expand Down Expand Up @@ -76,7 +79,9 @@ def deserialize(data: str) -> Any:
return json.loads(data)


def deprecated_argument(argument: str, replacement: str | None = None) -> Callable:
def deprecated_argument(
argument: str, replacement: str | None = None
) -> Callable[[Callable[P, R]], Callable[P, R]]:
"""
Decorator to warn if a deprecated argument is passed.

Expand All @@ -100,13 +105,16 @@ def test_method(cls, old_arg=None, new_arg=None):
if replacement:
message += f" Use {replacement} instead."

def decorator(func):
# Check if the function is a classmethod or staticmethod
def decorator(func: Callable[P, R]) -> Callable[P, R]:
# The documented usage puts this decorator inside
# @classmethod/@staticmethod, so this branch only covers the reversed
# order; a classmethod object is not a Callable to the type system,
# hence the casts.
if isinstance(func, (classmethod, staticmethod)):
underlying = func.__func__
underlying = cast(Callable[P, R], func.__func__)

@wraps(underlying)
def inner_wrapped(*args, **kwargs):
def inner_wrapped(*args: P.args, **kwargs: P.kwargs) -> R:
if argument in kwargs:
warn(message, DeprecationWarning, stacklevel=2)
else:
Expand All @@ -117,13 +125,13 @@ def inner_wrapped(*args, **kwargs):
return underlying(*args, **kwargs)

if isinstance(func, classmethod):
return classmethod(inner_wrapped)
return cast(Callable[P, R], classmethod(inner_wrapped))
else:
return staticmethod(inner_wrapped)
return cast(Callable[P, R], staticmethod(inner_wrapped))
else:

@wraps(func)
def inner_normal(*args, **kwargs):
def inner_normal(*args: P.args, **kwargs: P.kwargs) -> R:
if argument in kwargs:
warn(message, DeprecationWarning, stacklevel=2)
else:
Expand All @@ -148,15 +156,17 @@ def assert_no_warnings():
yield


def deprecated_function(name: str | None = None, replacement: str | None = None):
def deprecated_function(
name: str | None = None, replacement: str | None = None
) -> Callable[[Callable[P, R]], Callable[P, R]]:
"""
Decorator to mark a function as deprecated.

When the wrapped function is called, the decorator will log a deprecation
warning.
"""

def decorator(func):
def decorator(func: Callable[P, R]) -> Callable[P, R]:
fn_name = name or func.__name__
warning_message = (
f"Function {fn_name} is deprecated and will be "
Expand All @@ -166,7 +176,7 @@ def decorator(func):
warning_message += replacement

@wraps(func)
def wrapper(*args, **kwargs):
def wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
warn(warning_message, category=DeprecationWarning, stacklevel=3)
return func(*args, **kwargs)

Expand All @@ -175,7 +185,9 @@ def wrapper(*args, **kwargs):
return decorator


def deprecated_class(name: str | None = None, replacement: str | None = None):
def deprecated_class(
name: str | None = None, replacement: str | None = None
) -> Callable[[C], C]:
"""
Decorator to mark a class as deprecated.

Expand All @@ -193,7 +205,7 @@ class OldClass:
pass
"""

def decorator(cls):
def decorator(cls: C) -> C:
class_name = name or cls.__name__
warning_message = (
f"Class {class_name} is deprecated and will be "
Expand All @@ -202,10 +214,12 @@ def decorator(cls):
if replacement:
warning_message += replacement

original_init = cls.__init__
# getattr/setattr rather than attribute access: `cls` is typed as a
# class object, so mypy resolves `cls.__init__` to the metaclass slot.
original_init = getattr(cls, "__init__")

@wraps(original_init)
def new_init(self, *args, **kwargs):
def new_init(self: Any, *args: Any, **kwargs: Any) -> None:
# Emit only once per instance. When a deprecated subclass wraps a
# deprecated parent, both __init__ wrappers run via super().__init__;
# the sentinel keeps that to a single warning.
Expand All @@ -217,7 +231,7 @@ def new_init(self, *args, **kwargs):
pass
original_init(self, *args, **kwargs)

cls.__init__ = new_init
setattr(cls, "__init__", new_init)
return cls

return decorator
Expand Down
35 changes: 13 additions & 22 deletions redisvl/utils/vectorize/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -184,7 +184,7 @@ def embed_many(
if cache_misses:
cache_metadata = kwargs.pop("metadata", {})
new_embeddings = self._embed_many(
contents=cache_misses, batch_size=batch_size, **kwargs
cache_misses, batch_size=batch_size, **kwargs
)

# Store new embeddings in cache
Expand Down Expand Up @@ -317,7 +317,7 @@ async def aembed_many(
if cache_misses:
cache_metadata = kwargs.pop("metadata", {})
new_embeddings = await self._aembed_many(
contents=cache_misses, batch_size=batch_size, **kwargs
cache_misses, batch_size=batch_size, **kwargs
)

# Store new embeddings in cache
Expand All @@ -332,45 +332,36 @@ async def aembed_many(
# Process and return results
return [self._process_embedding(emb, as_buffer, self.dtype) for emb in results]

@deprecated_argument("text", "content")
def _embed(self, text: Any = "", content: Any = "", **kwargs) -> list[float]:
# The four hooks below are the provider extension points. Every caller —
# the public embed/embed_many/aembed/aembed_many above, and each provider's
# own _set_model_dims probe — passes the content as the first positional
# argument. None passes the deprecated `text`/`texts` alias, because the
# public methods resolve it before dispatching here.
def _embed(self, content: Any, **kwargs) -> list[float]:
"""Generate a vector embedding for a single item."""
raise NotImplementedError

@deprecated_argument("texts", "contents")
def _embed_many(
self,
contents: list[Any] | None = None,
texts: list[Any] | None = None,
batch_size: int = 10,
**kwargs,
self, contents: list[Any], batch_size: int = 10, **kwargs
) -> list[list[float]]:
"""Generate vector embeddings for a batch of items."""
raise NotImplementedError

@deprecated_argument("text", "content")
async def _aembed(self, content: Any = "", text: Any = "", **kwargs) -> list[float]:
async def _aembed(self, content: Any, **kwargs) -> list[float]:
"""Asynchronously generate a vector embedding for a single item."""
logger.warning(
"This vectorizer has no async embed method. Falling back to sync."
)
return self._embed(content=content or text, **kwargs)
return self._embed(content, **kwargs)

@deprecated_argument("texts", "contents")
async def _aembed_many(
self,
contents: list[Any] | None = None,
texts: list[Any] | None = None,
batch_size: int = 10,
**kwargs,
self, contents: list[Any], batch_size: int = 10, **kwargs
) -> list[list[float]]:
"""Asynchronously generate vector embeddings for a batch of items."""
logger.warning(
"This vectorizer has no async embed_many method. Falling back to sync."
)
return self._embed_many(
contents=contents or texts, batch_size=batch_size, **kwargs
)
return self._embed_many(contents, batch_size=batch_size, **kwargs)

def _get_from_cache_batch(
self, contents: list[Any], skip_cache: bool
Expand Down
25 changes: 5 additions & 20 deletions redisvl/utils/vectorize/text/azureopenai.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,6 @@
if TYPE_CHECKING:
from redisvl.extensions.cache.embeddings.embeddings import EmbeddingsCache

from redisvl.utils.utils import deprecated_argument
from redisvl.utils.vectorize.base import BaseVectorizer

# ignore that openai isn't imported
Expand Down Expand Up @@ -206,30 +205,27 @@ def _set_model_dims(self) -> int:
# fall back (TODO get more specific)
raise ValueError(f"Error setting embedding model dimensions: {str(e)}")

@deprecated_argument("text", "content")
@retry(
wait=wait_random_exponential(min=1, max=60),
stop=stop_after_attempt(6),
retry=retry_if_not_exception_type(TypeError),
reraise=True,
)
def _embed(self, content: str = "", text: str = "", **kwargs) -> list[float]:
def _embed(self, content: str, **kwargs) -> list[float]:
"""
Generate a vector embedding for a single text using the AzureOpenAI API.

Args:
content: Text to embed
text: Text to embed (deprecated - use `content` instead)
**kwargs: Additional parameters to pass to the AzureOpenAI API

Returns:
List[float]: Vector embedding as a list of floats

Raises:
TypeError: If text is not a string
TypeError: If content is not a string
ValueError: If embedding fails
"""
content = content or text
if not isinstance(content, str):
raise TypeError("Must pass in a str value to embed.")

Expand All @@ -241,7 +237,6 @@ def _embed(self, content: str = "", text: str = "", **kwargs) -> list[float]:
except Exception as e:
raise ValueError(f"Embedding text failed: {e}")

@deprecated_argument("texts", "contents")
@retry(
wait=wait_random_exponential(min=1, max=60),
stop=stop_after_attempt(6),
Expand All @@ -250,8 +245,7 @@ def _embed(self, content: str = "", text: str = "", **kwargs) -> list[float]:
)
def _embed_many(
self,
contents: list[str] | None = None,
texts: list[str] | None = None,
contents: list[str],
batch_size: int = 10,
**kwargs,
) -> list[list[float]]:
Expand All @@ -260,7 +254,6 @@ def _embed_many(

Args:
contents: List of texts to embed
texts: List of texts to embed (deprecated - use `contents` instead)
batch_size: Number of texts to process in each API call
**kwargs: Additional parameters to pass to the AzureOpenAI API

Expand All @@ -271,7 +264,6 @@ def _embed_many(
TypeError: If contents is not a list of strings
ValueError: If embedding fails
"""
contents = contents or texts
if not isinstance(contents, list):
raise TypeError("Must pass in a list of str values to embed.")
if contents and not isinstance(contents[0], str):
Expand All @@ -288,20 +280,18 @@ def _embed_many(
except Exception as e:
raise ValueError(f"Embedding texts failed: {e}")

@deprecated_argument("text", "content")
@retry(
wait=wait_random_exponential(min=1, max=60),
stop=stop_after_attempt(6),
retry=retry_if_not_exception_type(TypeError),
reraise=True,
)
async def _aembed(self, content: str = "", text: str = "", **kwargs) -> list[float]:
async def _aembed(self, content: str, **kwargs) -> list[float]:
"""
Asynchronously generate a vector embedding for a single text using the AzureOpenAI API.

Args:
content: Text to embed
text: Text to embed (deprecated - use `content` instead)
**kwargs: Additional parameters to pass to the AzureOpenAI API

Returns:
Expand All @@ -311,7 +301,6 @@ async def _aembed(self, content: str = "", text: str = "", **kwargs) -> list[flo
TypeError: If content is not a string
ValueError: If embedding fails
"""
content = content or text
if not isinstance(content, str):
raise TypeError("Must pass in a str value to embed.")

Expand All @@ -323,7 +312,6 @@ async def _aembed(self, content: str = "", text: str = "", **kwargs) -> list[flo
except Exception as e:
raise ValueError(f"Embedding text failed: {e}")

@deprecated_argument("texts", "contents")
@retry(
wait=wait_random_exponential(min=1, max=60),
stop=stop_after_attempt(6),
Expand All @@ -332,8 +320,7 @@ async def _aembed(self, content: str = "", text: str = "", **kwargs) -> list[flo
)
async def _aembed_many(
self,
contents: list[str] | None = None,
texts: list[str] | None = None,
contents: list[str],
batch_size: int = 10,
**kwargs,
) -> list[list[float]]:
Expand All @@ -342,7 +329,6 @@ async def _aembed_many(

Args:
contents: List of texts to embed
texts: List of texts to embed (deprecated - use `contents` instead)
batch_size: Number of texts to process in each API call
**kwargs: Additional parameters to pass to the AzureOpenAI API

Expand All @@ -353,7 +339,6 @@ async def _aembed_many(
TypeError: If contents is not a list of strings
ValueError: If embedding fails
"""
contents = contents or texts
if not isinstance(contents, list):
raise TypeError("Must pass in a list of str values to embed.")
if contents and not isinstance(contents[0], str):
Expand Down
31 changes: 5 additions & 26 deletions redisvl/utils/vectorize/text/bedrock.py
Original file line number Diff line number Diff line change
@@ -1,34 +1,13 @@
from typing import Any

from redisvl.utils.utils import deprecated_argument, deprecated_class
from redisvl.utils.utils import deprecated_class
from redisvl.utils.vectorize.bedrock import BedrockVectorizer


@deprecated_class(
name="BedrockTextVectorizer", replacement="Use BedrockVectorizer instead."
)
class BedrockTextVectorizer(BedrockVectorizer):
"""A backwards-compatible alias for BedrockVectorizer."""

@deprecated_argument("text", "content")
def embed(self, content: Any = "", text: Any = "", **kwargs) -> list[float] | bytes:
"""Generate a vector embedding for a single input using the AWS Bedrock API.

Deprecated: Use `BedrockVectorizer.embed` instead.
"""
content = content or text
return super().embed(content=content, **kwargs)

@deprecated_argument("texts", "contents")
def embed_many(
self,
contents: list[Any] | None = None,
texts: list[Any] | None = None,
**kwargs,
) -> list[list[float]]:
"""Generate vector embeddings for a batch of inputs using the AWS Bedrock API.
"""A backwards-compatible alias for BedrockVectorizer.

Deprecated: Use `BedrockVectorizer.embed_many` instead.
"""
contents = contents or texts
return super().embed_many(contents=contents, **kwargs)
The `text`/`texts` keyword arguments still work, and still warn, via
BaseVectorizer.
"""
Loading
Loading