302 lines
11 KiB
Python
302 lines
11 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import base64
|
|
from collections.abc import Sequence
|
|
from io import BytesIO
|
|
from typing import Any
|
|
|
|
import httpx
|
|
from PIL import Image
|
|
|
|
from cmvr_edge_ai.application import (
|
|
create_default_capability_registry,
|
|
create_default_registry,
|
|
)
|
|
from cmvr_edge_ai.config import load_config_data
|
|
from cmvr_edge_ai.contracts import BoundingBox, Detection, ImageFrame
|
|
from cmvr_edge_ai.detection import (
|
|
DetectionModel,
|
|
DetectionModelRegistry,
|
|
DetectionModelSpec,
|
|
)
|
|
from cmvr_edge_ai.server.application import EdgeAIServerApplication
|
|
|
|
|
|
class _MobilePhoneModel(DetectionModel):
|
|
def __init__(self) -> None:
|
|
self.loaded = False
|
|
self.closed = False
|
|
self.predict_calls = 0
|
|
|
|
def load(self) -> None:
|
|
self.loaded = True
|
|
|
|
def predict(
|
|
self,
|
|
frame: ImageFrame,
|
|
labels: tuple[str, ...],
|
|
confidence: float,
|
|
) -> Sequence[Detection]:
|
|
assert self.loaded
|
|
assert frame.pixel_format == "BGR8"
|
|
assert labels == ("mobile_phone",)
|
|
assert confidence == 0.5
|
|
self.predict_calls += 1
|
|
return (
|
|
Detection(
|
|
label="mobile_phone",
|
|
confidence=0.91,
|
|
box=BoundingBox(1, 1, frame.width - 1, frame.height - 1),
|
|
),
|
|
)
|
|
|
|
def close(self) -> None:
|
|
self.closed = True
|
|
|
|
|
|
def _jpeg() -> bytes:
|
|
image = Image.new("RGB", (8, 6), (30, 60, 90))
|
|
output = BytesIO()
|
|
image.save(output, format="JPEG")
|
|
return output.getvalue()
|
|
|
|
|
|
async def _keep_event_loop_responsive() -> None:
|
|
# Restricted CI containers may not let a ThreadPoolExecutor helper wake the
|
|
# event loop. A small timer mirrors the workaround used by the existing
|
|
# detection integration tests and does not change production behavior.
|
|
while True:
|
|
await asyncio.sleep(0.01)
|
|
|
|
|
|
def _config(): # type: ignore[no-untyped-def]
|
|
return load_config_data(
|
|
{
|
|
"api_version": "cmvr.edge.ai/v1",
|
|
"runtime": {"thread_workers": 1},
|
|
"server": {
|
|
"enabled": True,
|
|
"http": {"request_timeout_s": 5},
|
|
"routes": {
|
|
"detect.mobile_phone": {
|
|
"pipeline": "mobile_remote",
|
|
"model_id": "yolov8n-mobile-phone@1",
|
|
"queue_capacity": 2,
|
|
}
|
|
},
|
|
},
|
|
"pipelines": {
|
|
"mobile_remote": {
|
|
"nodes": {
|
|
"requests": {
|
|
"uses": "server.request_source@1",
|
|
"with": {"category": "detect.mobile_phone"},
|
|
},
|
|
"decode": {"uses": "media.image_decoder.pillow@1"},
|
|
"detector": {
|
|
"uses": "detection.model@1",
|
|
"with": {
|
|
"model": "yolov8n-mobile-phone@1",
|
|
"confidence": 0.5,
|
|
"attach_frame": True,
|
|
},
|
|
},
|
|
"response": {
|
|
"uses": "server.detection_response@1",
|
|
"with": {
|
|
"category": "detect.mobile_phone",
|
|
"backend": "test",
|
|
},
|
|
},
|
|
"responses": {"uses": "server.response_sink@1"},
|
|
},
|
|
"edges": [
|
|
{"from": "requests.requests", "to": "decode.requests"},
|
|
{"from": "decode.frames", "to": "detector.frames"},
|
|
{
|
|
"from": "detector.detections",
|
|
"to": "response.detections",
|
|
},
|
|
{
|
|
"from": "response.responses",
|
|
"to": "responses.responses",
|
|
},
|
|
],
|
|
}
|
|
},
|
|
}
|
|
)
|
|
|
|
|
|
def test_remote_detection_runs_through_the_real_pipeline_runtime() -> None:
|
|
async def exercise() -> tuple[dict[str, Any], _MobilePhoneModel, bool]:
|
|
model = _MobilePhoneModel()
|
|
model_registry = DetectionModelRegistry()
|
|
model_registry.register(
|
|
DetectionModelSpec(
|
|
model_id="yolov8n-mobile-phone@1",
|
|
name="Test Mobile Phone",
|
|
supported_labels=("mobile_phone",),
|
|
factory=lambda options: model,
|
|
backend="test",
|
|
)
|
|
)
|
|
app = EdgeAIServerApplication(
|
|
_config(),
|
|
create_default_registry(
|
|
discover_entry_points=False,
|
|
model_registry=model_registry,
|
|
),
|
|
create_default_capability_registry(discover_entry_points=False),
|
|
)
|
|
await app.start()
|
|
ticker = asyncio.create_task(_keep_event_loop_responsive())
|
|
try:
|
|
async with httpx.AsyncClient(
|
|
transport=httpx.ASGITransport(app=app.asgi_app),
|
|
base_url="http://test",
|
|
) as client:
|
|
response = await client.post(
|
|
"/v1/inference",
|
|
json={
|
|
"schema_version": "cmvr.inference-request/v1",
|
|
"request_id": "remote-1",
|
|
"category": "detect.mobile_phone",
|
|
"source_id": "edge-wlan-192.168.1.20",
|
|
"inputs": [
|
|
{
|
|
"kind": "image",
|
|
"name": "image",
|
|
"media_type": "image/jpeg",
|
|
"encoding": "base64",
|
|
"data": base64.b64encode(_jpeg()).decode("ascii"),
|
|
}
|
|
],
|
|
"requested_artifact_roles": ["annotated", "original"],
|
|
},
|
|
)
|
|
ready = await client.get("/health/ready")
|
|
return response.json(), model, ready.status_code == 200
|
|
finally:
|
|
await app.stop()
|
|
ticker.cancel()
|
|
await asyncio.gather(ticker, return_exceptions=True)
|
|
|
|
payload, model, ready = asyncio.run(exercise())
|
|
|
|
assert ready is True
|
|
assert payload["schema_version"] == "cmvr.inference-response/v1", payload
|
|
assert payload["request_id"] == "remote-1"
|
|
assert payload["category"] == "detect.mobile_phone"
|
|
assert payload["source_id"] == "edge-wlan-192.168.1.20"
|
|
assert payload["status"] == "succeeded"
|
|
assert payload["outputs"][0]["kind"] == "detections"
|
|
assert payload["outputs"][0]["items"][0]["label"] == "mobile_phone"
|
|
assert [item["role"] for item in payload["artifacts"]] == [
|
|
"annotated",
|
|
"original",
|
|
]
|
|
assert all(
|
|
base64.b64decode(item["data"]).startswith(b"\xff\xd8")
|
|
for item in payload["artifacts"]
|
|
)
|
|
assert model.predict_calls == 1
|
|
assert model.closed is True
|
|
|
|
|
|
def test_malformed_image_fails_only_its_request_and_pipeline_stays_ready() -> None:
|
|
async def exercise() -> tuple[int, int, dict[str, Any], int, bool]:
|
|
model = _MobilePhoneModel()
|
|
model_registry = DetectionModelRegistry()
|
|
model_registry.register(
|
|
DetectionModelSpec(
|
|
model_id="yolov8n-mobile-phone@1",
|
|
name="Test Mobile Phone",
|
|
supported_labels=("mobile_phone",),
|
|
factory=lambda options: model,
|
|
backend="test",
|
|
)
|
|
)
|
|
app = EdgeAIServerApplication(
|
|
_config(),
|
|
create_default_registry(
|
|
discover_entry_points=False,
|
|
model_registry=model_registry,
|
|
),
|
|
create_default_capability_registry(discover_entry_points=False),
|
|
)
|
|
await app.start()
|
|
ticker = asyncio.create_task(_keep_event_loop_responsive())
|
|
try:
|
|
async with httpx.AsyncClient(
|
|
transport=httpx.ASGITransport(app=app.asgi_app),
|
|
base_url="http://test",
|
|
) as client:
|
|
def payload(request_id: str, image: bytes) -> dict[str, Any]:
|
|
return {
|
|
"schema_version": "cmvr.inference-request/v1",
|
|
"request_id": request_id,
|
|
"category": "detect.mobile_phone",
|
|
"inputs": [
|
|
{
|
|
"kind": "image",
|
|
"name": "image",
|
|
"media_type": "image/jpeg",
|
|
"encoding": "base64",
|
|
"data": base64.b64encode(image).decode("ascii"),
|
|
}
|
|
],
|
|
}
|
|
|
|
malformed = await client.post(
|
|
"/v1/inference",
|
|
json=payload("bad-jpeg", b"not-a-jpeg"),
|
|
)
|
|
valid = await client.post(
|
|
"/v1/inference",
|
|
json=payload("good-jpeg", _jpeg()),
|
|
)
|
|
ready = await client.get("/health/ready")
|
|
return (
|
|
malformed.status_code,
|
|
valid.status_code,
|
|
valid.json(),
|
|
model.predict_calls,
|
|
ready.status_code == 200,
|
|
)
|
|
finally:
|
|
await app.stop()
|
|
ticker.cancel()
|
|
await asyncio.gather(ticker, return_exceptions=True)
|
|
|
|
bad_status, good_status, good_payload, predict_calls, ready = asyncio.run(
|
|
exercise()
|
|
)
|
|
|
|
assert bad_status == 500
|
|
assert good_status == 200
|
|
assert good_payload["request_id"] == "good-jpeg"
|
|
assert good_payload["status"] == "succeeded"
|
|
assert predict_calls == 1
|
|
assert ready is True
|
|
|
|
|
|
def test_server_application_rejects_category_model_mismatch_before_start() -> None:
|
|
config_data = _config().model_dump(by_alias=True)
|
|
config_data["server"]["routes"]["detect.phone_use"] = config_data["server"][
|
|
"routes"
|
|
].pop("detect.mobile_phone")
|
|
config = load_config_data(config_data)
|
|
|
|
try:
|
|
EdgeAIServerApplication(
|
|
config,
|
|
create_default_registry(discover_entry_points=False),
|
|
create_default_capability_registry(discover_entry_points=False),
|
|
)
|
|
except ValueError as exc:
|
|
assert "does not match capability category" in str(exc)
|
|
else:
|
|
raise AssertionError("category/model mismatch must fail during construction")
|