"""CPU regression checks for image forwarding, probabilities and input errors.""" import base64 import io import json import unittest import httpx from PIL import Image import decisions_server as server from decisions_api import ApiError, assemble, build_plan from image_inputs import image_input from image_jevbench import OpenJevImageAdapter def picture(): b = io.BytesIO() Image.new("RGB", (32, 32), "red").save(b, format="PNG") return b.getvalue() class Images(unittest.IsolatedAsyncioTestCase): async def asyncSetUp(self): self.calls = [] def upstream(req): body = json.loads(req.content) self.calls.append(body) return httpx.Response(200, json=[{"embedding": [0, 2, -1], "meta_info": {"prompt_tokens": 80}} for _ in body["text"]]) server._client = httpx.AsyncClient(transport=httpx.MockTransport(upstream)) self.api = httpx.AsyncClient(transport=httpx.ASGITransport(app=server.app), base_url="http://test") server.API_KEY = "" self.body = {"model": "test", "state": "Image: <>", "image_data": base64.b64encode(picture()).decode(), "questions": {"q": {"type": "choice", "instructions": "Which color?", "criteria": {"a": "red", "b": "blue"}}}} async def asyncTearDown(self): await self.api.aclose() await server._client.aclose() server._client = None async def test_pixels_forwarded_per_option(self): r = await self.api.post("/v1/systemone", json=self.body) self.assertEqual(r.status_code, 200, r.text) self.assertEqual(r.json()["answers"]["q"]["probabilities"], {"a": .5, "b": .5}) self.assertEqual(len(self.calls), 1) for c in self.calls: self.assertEqual(c["image_data"], [self.body["image_data"]] * 2) self.assertEqual(c["text"][0].count("<|image_pad|>"), 1) self.assertNotIn("<>", c["text"][0]) async def test_text_request_still_batched(self): del self.body["image_data"] self.body["state"] = "A red box." r = await self.api.post("/v1/systemone", json=self.body) self.assertEqual(r.status_code, 200) self.assertEqual(len(self.calls), 1) self.assertNotIn("image_data", self.calls[0]) async def test_bad_image_does_not_reach_worker(self): for data in ("/etc/passwd", "https://example.com/image.png", "not-base64", [], ""): self.body["image_data"] = data r = await self.api.post("/v1/systemone", json=self.body) self.assertEqual(r.status_code, 400, r.text) self.assertEqual(self.calls, []) async def test_duplicate_marker_rejected(self): self.body["state"] = "<> <>" r = await self.api.post("/v1/systemone", json=self.body) self.assertEqual(r.status_code, 400) self.assertEqual(self.calls, []) def test_data_uri_and_zero_entailment(self): self.assertEqual(image_input({"image_data": "data:image/png;base64," + self.body["image_data"]}), self.body["image_data"]) p = assemble(build_plan(self.body), [0, 0])["q"]["probabilities"] self.assertEqual(p, {"a": .5, "b": .5}) with self.assertRaises(ValueError): assemble(build_plan(self.body), [float("nan"), 1]) def test_benchmark_does_not_send_gold_or_alt(self): a = OpenJevImageAdapter("http://test", "test") try: body = a.build_request({"question": "Q", "options": [{"label": "a", "text": "red"}], "correctLabel": "SECRET_GOLD", "alt": "SECRET_DESCRIPTION"}, picture()) self.assertNotIn("SECRET", json.dumps(body)) finally: a.close() if __name__ == "__main__": unittest.main()