GET api update for image vector
This commit is contained in:
@@ -93,6 +93,52 @@ def test_unsearchable_vectors_are_refused_with_a_reason(bad, msg):
|
||||
im.normalise_vector(bad)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# the query-string forms of the vector (GET /search/image-vector)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_base64_round_trips_to_float32_precision():
|
||||
v = _unit()
|
||||
encoded = im.encode_vector_b64(v)
|
||||
|
||||
assert len(encoded) < 5_500 and "=" not in encoded # 5,462 chars, unpadded
|
||||
back = im.parse_vector_param(encoded)
|
||||
assert len(back) == 1024 and max(abs(a - b) for a, b in zip(v, back)) < 1e-6
|
||||
|
||||
|
||||
def test_the_standard_base64_alphabet_and_padding_are_accepted_too():
|
||||
import base64, struct
|
||||
v = _unit()
|
||||
standard = base64.b64encode(struct.pack("<1024f", *v)).decode() # '+', '/', '=' padding
|
||||
|
||||
assert max(abs(a - b) for a, b in zip(v, im.parse_vector_param(standard))) < 1e-6
|
||||
|
||||
|
||||
def test_comma_separated_values_tolerate_brackets_whitespace_and_newlines():
|
||||
v = _unit()
|
||||
text = "[ " + ("," + chr(10) + " ").join(f"{x:.6f}" for x in v) + " ]"
|
||||
|
||||
back = im.parse_vector_param(text)
|
||||
assert len(back) == 1024 and max(abs(a - b) for a, b in zip(v, back)) < 1e-6
|
||||
|
||||
|
||||
@pytest.mark.parametrize("raw, msg", [
|
||||
("", "empty"),
|
||||
(" ", "empty"),
|
||||
("0.1, 0.2, x, 0.4", "value 3"),
|
||||
("MTIz", "decodes to 3 bytes"),
|
||||
("!!!not base64!!!", "comma-separated numbers or base64"),
|
||||
])
|
||||
def test_unparseable_query_vectors_say_what_is_wrong(raw, msg):
|
||||
with pytest.raises(im.InvalidVectorError, match=msg):
|
||||
im.parse_vector_param(raw)
|
||||
|
||||
|
||||
def test_a_short_csv_vector_is_caught_by_the_same_length_rule_as_the_post_body():
|
||||
with pytest.raises(im.InvalidVectorError, match="1024 values"):
|
||||
im.normalise_vector(im.parse_vector_param(",".join(["0.1"] * 1023)))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# the label text
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -116,6 +116,57 @@ def test_a_service_rejection_is_422_not_500(client, monkeypatch):
|
||||
assert res.status_code == 422 and "all zeros" in res.text
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /search/image-vector
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_get_with_a_base64_vector_returns_the_same_cards_and_stays_under_8kb(client, fake_search):
|
||||
from app.services.image_match import encode_vector_b64
|
||||
v = _unit()
|
||||
params = {"vector": encode_vector_b64(v), "text": "Britannia Marie Gold 300 g", "top_k": 5, "min_score": 0.1}
|
||||
|
||||
res = client.get("/api/search/image-vector", params=params)
|
||||
|
||||
assert res.status_code == 200, res.text
|
||||
assert res.json()["results"][0]["product_name"] == "Britannia Marie Gold 300g"
|
||||
call = fake_search[0]
|
||||
assert len(call["vector"]) == 1024 and max(abs(a - b) for a, b in zip(v, call["vector"])) < 1e-6
|
||||
assert call["text"] == "Britannia Marie Gold 300 g" and call["top_k"] == 5 and call["min_score"] == 0.1
|
||||
# the nginx on the app domain allows an 8 KB request line; base64 must fit it
|
||||
assert len(str(res.request.url)) < 8_000
|
||||
|
||||
|
||||
def test_get_with_comma_separated_values_works_on_the_api_host(client, fake_search):
|
||||
v = _unit()
|
||||
res = client.get("/api/search/image-vector", params={"vector": ",".join(f"{x:.5f}" for x in v)})
|
||||
|
||||
assert res.status_code == 200, res.text
|
||||
assert len(fake_search[0]["vector"]) == 1024
|
||||
assert max(abs(a - b) for a, b in zip(v, fake_search[0]["vector"])) < 1e-5
|
||||
|
||||
|
||||
@pytest.mark.parametrize("params, fragment", [
|
||||
({}, "vector"),
|
||||
({"vector": "0.1,0.2,x"}, "value 3"),
|
||||
({"vector": "MTIz"}, "decodes to 3 bytes"),
|
||||
({"vector": ",".join(["0.1"] * 1023)}, "1024 values"),
|
||||
({"vector": ",".join(["0.0"] * 1024)}, "all zeros"),
|
||||
({"vector": ",".join(["0.1"] * 1024), "top_k": 51}, "top_k"),
|
||||
])
|
||||
def test_get_rejects_bad_input_with_422_and_a_reason(client, fake_search, params, fragment):
|
||||
res = client.get("/api/search/image-vector", params=params)
|
||||
assert res.status_code == 422 and fragment in res.text
|
||||
assert fake_search == []
|
||||
|
||||
|
||||
def test_get_and_post_agree(client, fake_search):
|
||||
from app.services.image_match import encode_vector_b64
|
||||
v = _unit()
|
||||
a = client.get("/api/search/image-vector", params={"vector": encode_vector_b64(v), "top_k": 3}).json()
|
||||
b = client.post("/api/search/image-vector", json={"vector": v, "top_k": 3}).json()
|
||||
assert [r["image_id"] for r in a["results"]] == [r["image_id"] for r in b["results"]]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /search/image
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -172,3 +223,4 @@ def test_both_routes_are_documented(client):
|
||||
paths = client.get("/openapi.json").json()["paths"]
|
||||
assert "/api/search/image-vector" in paths and "/api/search/image" in paths
|
||||
assert "post" in paths["/api/search/image"] and "get" in paths["/api/search"]
|
||||
assert {"get", "post"} <= set(paths["/api/search/image-vector"])
|
||||
|
||||
Reference in New Issue
Block a user