codekingpro/portable-devtools
114k
1import shutil
2import os
3from typing import List, Hashable
4
5import hypothesis.strategies as st
6import onnxruntime
7import pytest
8from hypothesis import given, settings
9
10from chromadb.utils.embedding_functions.onnx_mini_lm_l6_v2 import (
11 ONNXMiniLM_L6_V2,
12)
13
14from chromadb.utils.embedding_functions.onnx_mini_lm_l6_v2 import _verify_sha256
15
16
17def unique_by(x: Hashable) -> Hashable:
18 return x
19
20
21@settings(deadline=None)
22@given(
23 providers=st.lists(
24 st.sampled_from(onnxruntime.get_all_providers()).filter(
25 lambda x: x not in onnxruntime.get_available_providers()
26 ),
27 unique_by=unique_by,
28 min_size=1,
29 )
30)
31def test_unavailable_provider_multiple(providers: List[str]) -> None:
32 with pytest.raises(ValueError) as e:
33 ef = ONNXMiniLM_L6_V2(preferred_providers=providers)
34 ef(["test"])
35 assert "Preferred providers must be subset of available providers" in str(e.value)
36
37
38@given(
39 providers=st.lists(
40 st.sampled_from(onnxruntime.get_available_providers()),
41 min_size=1,
42 unique_by=unique_by,
43 )
44)
45def test_available_provider(providers: List[str]) -> None:
46 ef = ONNXMiniLM_L6_V2(preferred_providers=providers)
47 ef(["test"])
48
49
50def test_warning_no_providers_supplied() -> None:
51 ef = ONNXMiniLM_L6_V2()
52 ef(["test"])
53
54
55@given(
56 providers=st.lists(
57 st.sampled_from(onnxruntime.get_available_providers()),
58 min_size=1,
59 ).filter(lambda x: len(x) > len(set(x)))
60)
61def test_provider_repeating(providers: List[str]) -> None:
62 with pytest.raises(ValueError) as e:
63 ef = ONNXMiniLM_L6_V2(preferred_providers=providers)
64 ef(["test"])
65 assert "Preferred providers must be unique" in str(e.value)
66
67
68def test_invalid_sha256() -> None:
69 ef = ONNXMiniLM_L6_V2()
70 shutil.rmtree(ef.DOWNLOAD_PATH) # clean up any existing models
71 with pytest.raises(ValueError) as e:
72 ef._MODEL_SHA256 = "invalid"
73 ef(["test"])
74 assert "does not match expected SHA256 hash" in str(e.value)
75
76
77def test_partial_download() -> None:
78 ef = ONNXMiniLM_L6_V2()
79 shutil.rmtree(ef.DOWNLOAD_PATH, ignore_errors=True) # clean up any existing models
80 os.makedirs(ef.DOWNLOAD_PATH, exist_ok=True)
81 path = os.path.join(ef.DOWNLOAD_PATH, ef.ARCHIVE_FILENAME)
82 with open(path, "wb") as f: # create invalid file to simulate partial download
83 f.write(b"invalid")
84 ef._download_model_if_not_exists() # re-download model
85 assert os.path.exists(path)
86 assert _verify_sha256(
87 str(os.path.join(ef.DOWNLOAD_PATH, ef.ARCHIVE_FILENAME)),
88 ef._MODEL_SHA256,
89 )
90 assert len(ef(["test"])) == 1
91 