import io import zipfile import pytest from fastapi import FastAPI from fastapi.testclient import TestClient from PIL import Image def make_png_bytes(color=(255, 0, 0, 255), size=(8, 8)): img = Image.new("RGBA", size, color) buf = io.BytesIO() img.save(buf, format="PNG") return buf.getvalue() class FakeInf: def __init__(self): self.killed = 0 async def remove_bg(self, data, label=None): return make_png_bytes(color=(0, 255, 0, 128)) def queue_depth(self): return 0 def in_flight(self): return 0 def jobs(self): return [ {"id": "abc", "status": "processing", "label": "a.png", "age_seconds": 1.2}, ] def kill_all(self): self.killed += 1 return {"killed_pending": 0, "worker_restarted": True} @pytest.fixture def client(monkeypatch): import api fake = FakeInf() monkeypatch.setattr(api, "inf", fake) app = FastAPI() app.include_router(api.router) return TestClient(app) def test_health(client): r = client.get("/health") assert r.status_code == 200 body = r.json() assert body["status"] == "ok" assert body["model"] == "birefnet-general" assert body["queue"] == 0 assert body["in_flight"] == 0 def test_jobs(client): r = client.get("/jobs") assert r.status_code == 200 body = r.json() assert body["total"] == 1 assert "jobs" in body assert body["jobs"][0]["status"] == "processing" assert body["jobs"][0]["label"] == "a.png" def test_remove_single(client): r = client.post( "/remove", files={"file": ("a.png", make_png_bytes(), "image/png")}, ) assert r.status_code == 200 assert r.headers["content-type"] == "image/png" assert r.content[:8] == b"\x89PNG\r\n\x1a\n" def test_remove_rejects_non_image(client): r = client.post( "/remove", files={"file": ("a.txt", b"hello", "text/plain")}, ) assert r.status_code == 400 def test_remove_batch_returns_zip(client): files = [ ("files", ("a.png", make_png_bytes(), "image/png")), ("files", ("b.png", make_png_bytes(), "image/png")), ] r = client.post("/remove/batch", files=files) assert r.status_code == 200 assert r.headers["content-type"] == "application/zip" zf = zipfile.ZipFile(io.BytesIO(r.content)) names = zf.namelist() assert "a.png" in names assert "b.png" in names def test_remove_batch_rejects_too_many(client, monkeypatch): import config monkeypatch.setattr(config, "MAX_BATCH_FILES", 1) files = [ ("files", ("a.png", make_png_bytes(), "image/png")), ("files", ("b.png", make_png_bytes(), "image/png")), ] r = client.post("/remove/batch", files=files) assert r.status_code == 400 def test_remove_batch_records_per_file_error(client, monkeypatch): import api calls = {"n": 0} orig = api.inf.remove_bg async def flaky(data, label=None): calls["n"] += 1 if calls["n"] == 1: raise RuntimeError("boom") return await orig(data, label=label) monkeypatch.setattr(api.inf, "remove_bg", flaky) files = [ ("files", ("bad.png", make_png_bytes(), "image/png")), ("files", ("good.png", make_png_bytes(), "image/png")), ] r = client.post("/remove/batch", files=files) assert r.status_code == 200 zf = zipfile.ZipFile(io.BytesIO(r.content)) assert "_errors.txt" in zf.namelist() assert "good.png" in zf.namelist() assert b"bad.png" in zf.read("_errors.txt") def test_kill_without_configured_key_returns_503(client, monkeypatch): import config monkeypatch.setattr(config, "KILL_API_KEY", None) r = client.get("/kill") assert r.status_code == 503 def test_kill_wrong_key_returns_401(client, monkeypatch): import config monkeypatch.setattr(config, "KILL_API_KEY", "secret") r = client.get("/kill", params={"key": "nope"}) assert r.status_code == 401 def test_kill_missing_key_returns_401(client, monkeypatch): import config monkeypatch.setattr(config, "KILL_API_KEY", "secret") r = client.get("/kill") assert r.status_code == 401 def test_kill_correct_key_succeeds(client, monkeypatch): import config monkeypatch.setattr(config, "KILL_API_KEY", "secret") r = client.get("/kill", params={"key": "secret"}) assert r.status_code == 200 body = r.json() assert body["worker_restarted"] is True