codekingpro/portable-devtools
114k
1from chromadb.api.types import EmbeddingFunction, Embeddable, Embeddings
2import numpy as np
3from typing import cast, Any
4from chromadb.utils.embedding_functions import (
5 register_embedding_function,
6 known_embedding_functions,
7)
8
9
10class LegacyCustomEmbeddingFunction(EmbeddingFunction[Embeddable]):
11 def __call__(self, input: Embeddable) -> Embeddings:
12 return cast(Embeddings, np.array([1, 2, 3]).tolist())
13
14
15class CustomEmbeddingFunction(EmbeddingFunction[Embeddable]):
16 def __call__(self, input: Embeddable) -> Embeddings:
17 return cast(Embeddings, np.array([1, 2, 3]).tolist())
18
19 def __init__(self, *args: Any, **kwargs: Any) -> None:
20 pass
21
22 @staticmethod
23 def name() -> str:
24 return "custom_embedding_function"
25
26 @staticmethod
27 def build_from_config(config: dict[str, Any]) -> "CustomEmbeddingFunction":
28 return CustomEmbeddingFunction()
29
30 def get_config(self) -> dict[str, Any]:
31 return {}
32
33
34@register_embedding_function
35class CustomEmbeddingFunctionWithRegistration(EmbeddingFunction[Embeddable]):
36 def __call__(self, input: Embeddable) -> Embeddings:
37 return cast(Embeddings, np.array([1, 2, 3]).tolist())
38
39 def __init__(self, *args: Any, **kwargs: Any) -> None:
40 pass
41
42 @staticmethod
43 def name() -> str:
44 return "custom_embedding_function_with_registration"
45
46 @staticmethod
47 def build_from_config(
48 config: dict[str, Any]
49 ) -> "CustomEmbeddingFunctionWithRegistration":
50 return CustomEmbeddingFunctionWithRegistration()
51
52 def get_config(self) -> dict[str, Any]:
53 return {}
54
55
56def test_legacy_custom_ef() -> None:
57 ef = LegacyCustomEmbeddingFunction()
58 result = ef(["test"])
59
60 # Check the structure: we expect a list with one NumPy array
61 assert isinstance(result, list), "Result should be a list"
62 assert len(result) == 1, "Result should contain exactly one element"
63 assert isinstance(result[0], np.ndarray), "Result element should be a NumPy array"
64
65 # Compare the contents of the array
66 expected = np.array([1, 2, 3], dtype=np.float32)
67 assert np.array_equal(
68 result[0], expected
69 ), f"Arrays not equal: {result[0]} vs {expected}"
70
71
72def test_custom_ef() -> None:
73 ef = CustomEmbeddingFunction()
74 result = ef(["test"])
75
76 # Same checks as above
77 assert isinstance(result, list), "Result should be a list"
78 assert len(result) == 1, "Result should contain exactly one element"
79 assert isinstance(result[0], np.ndarray), "Result element should be a NumPy array"
80
81 expected = np.array([1, 2, 3], dtype=np.float32)
82 assert np.array_equal(
83 result[0], expected
84 ), f"Arrays not equal: {result[0]} vs {expected}"
85
86
87def test_custom_ef_registration() -> None:
88 # check all 4 embedding functions for registration.
89 # LegacyCustomEmbeddingFunction should not be in known_embedding_functions
90 # CustomEmbeddingFunction should not be in known_embedding_functions
91 # CustomEmbeddingFunctionWithRegistration should be in known_embedding_functions
92
93 assert "legacy_custom_embedding_function" not in known_embedding_functions
94 assert "custom_embedding_function" not in known_embedding_functions
95 assert "custom_embedding_function_with_registration" in known_embedding_functions
96 