recipe: harden auth/metrics, add rate limiting and tests
This commit is contained in:
@@ -0,0 +1,223 @@
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import sqlite3
|
||||
import tempfile
|
||||
|
||||
import requests
|
||||
|
||||
os.environ.setdefault(
|
||||
"RECIPE_METRICS_DB",
|
||||
os.path.join(tempfile.gettempdir(), "recipe-service-tests-metrics.sqlite3"),
|
||||
)
|
||||
|
||||
from app import app
|
||||
|
||||
|
||||
class MockResponse:
|
||||
def __init__(self, status_code: int, payload: dict | None = None):
|
||||
self.status_code = status_code
|
||||
self._payload = payload or {}
|
||||
|
||||
def json(self) -> dict:
|
||||
return self._payload
|
||||
|
||||
|
||||
def auth_header() -> dict[str, str]:
|
||||
return {"Authorization": "Bearer test-token"}
|
||||
|
||||
|
||||
def make_image(content: bytes = b"img") -> tuple[io.BytesIO, str, str]:
|
||||
return io.BytesIO(content), "sample.png", "image/png"
|
||||
|
||||
|
||||
def test_health() -> None:
|
||||
client = app.test_client()
|
||||
response = client.get("/health")
|
||||
assert response.status_code == 200
|
||||
assert response.get_json() == {"status": "ok"}
|
||||
|
||||
|
||||
def test_requires_authorization(monkeypatch) -> None:
|
||||
monkeypatch.setenv("RECIPE_API_TOKEN", "test-token")
|
||||
client = app.test_client()
|
||||
response = client.post("/receipt")
|
||||
assert response.status_code == 401
|
||||
assert response.get_json() == {"error": "unauthorized"}
|
||||
|
||||
|
||||
def test_missing_image_or_prompt(monkeypatch) -> None:
|
||||
monkeypatch.setenv("RECIPE_API_TOKEN", "test-token")
|
||||
client = app.test_client()
|
||||
response = client.post("/receipt", headers=auth_header())
|
||||
assert response.status_code == 400
|
||||
assert response.get_json() == {"error": "image and prompt are required"}
|
||||
|
||||
|
||||
def test_unsupported_type(monkeypatch) -> None:
|
||||
monkeypatch.setenv("RECIPE_API_TOKEN", "test-token")
|
||||
client = app.test_client()
|
||||
response = client.post(
|
||||
"/receipt",
|
||||
headers=auth_header(),
|
||||
data={
|
||||
"prompt": "p",
|
||||
"image": (io.BytesIO(b"x"), "bad.txt", "text/plain"),
|
||||
},
|
||||
content_type="multipart/form-data",
|
||||
)
|
||||
assert response.status_code == 415
|
||||
assert response.get_json() == {"error": "unsupported image type"}
|
||||
|
||||
|
||||
def test_image_too_large(monkeypatch) -> None:
|
||||
monkeypatch.setenv("RECIPE_API_TOKEN", "test-token")
|
||||
client = app.test_client()
|
||||
payload = b"a" * (10 * 1024 * 1024 + 1)
|
||||
response = client.post(
|
||||
"/receipt",
|
||||
headers=auth_header(),
|
||||
data={
|
||||
"prompt": "p",
|
||||
"image": (io.BytesIO(payload), "big.png", "image/png"),
|
||||
},
|
||||
content_type="multipart/form-data",
|
||||
)
|
||||
assert response.status_code == 413
|
||||
assert response.get_json() == {"error": "image is too large"}
|
||||
|
||||
|
||||
def test_success(monkeypatch) -> None:
|
||||
monkeypatch.setenv("RECIPE_API_TOKEN", "test-token")
|
||||
|
||||
def fake_post(*args, **kwargs):
|
||||
return MockResponse(200, {"text": "ok", "usage": {"totalTokens": 10}})
|
||||
|
||||
monkeypatch.setattr(requests, "post", fake_post)
|
||||
client = app.test_client()
|
||||
response = client.post(
|
||||
"/receipt",
|
||||
headers=auth_header(),
|
||||
data={
|
||||
"prompt": "p",
|
||||
"image": make_image(),
|
||||
},
|
||||
content_type="multipart/form-data",
|
||||
)
|
||||
assert response.status_code == 200
|
||||
assert response.get_json() == {"text": "ok", "usage": {"totalTokens": 10}}
|
||||
|
||||
|
||||
def test_upstream_502_contract(monkeypatch) -> None:
|
||||
monkeypatch.setenv("RECIPE_API_TOKEN", "test-token")
|
||||
|
||||
def fake_post(*args, **kwargs):
|
||||
return MockResponse(500, {"error": {"message": "provider detail"}})
|
||||
|
||||
monkeypatch.setattr(requests, "post", fake_post)
|
||||
client = app.test_client()
|
||||
response = client.post(
|
||||
"/receipt",
|
||||
headers=auth_header(),
|
||||
data={
|
||||
"prompt": "p",
|
||||
"image": make_image(),
|
||||
},
|
||||
content_type="multipart/form-data",
|
||||
)
|
||||
assert response.status_code == 502
|
||||
assert response.get_json() == {
|
||||
"error": "upstream recognition failed",
|
||||
"code": "upstream_error",
|
||||
}
|
||||
|
||||
|
||||
def test_rate_limit_returns_429(monkeypatch) -> None:
|
||||
monkeypatch.setenv("RECIPE_API_TOKEN", "test-token")
|
||||
|
||||
def fake_count_since(client_ip: str, started_at_from: str) -> int:
|
||||
return 20
|
||||
|
||||
monkeypatch.setattr("app.count_since", fake_count_since)
|
||||
client = app.test_client()
|
||||
response = client.post(
|
||||
"/receipt",
|
||||
headers=auth_header(),
|
||||
data={
|
||||
"prompt": "p",
|
||||
"image": make_image(),
|
||||
},
|
||||
content_type="multipart/form-data",
|
||||
environ_base={"REMOTE_ADDR": "198.51.100.10"},
|
||||
)
|
||||
assert response.status_code == 429
|
||||
assert response.get_json() == {"error": "too many requests", "code": "rate_limited"}
|
||||
|
||||
|
||||
def test_rate_limit_allows_below_threshold(monkeypatch) -> None:
|
||||
monkeypatch.setenv("RECIPE_API_TOKEN", "test-token")
|
||||
|
||||
def fake_count_since(client_ip: str, started_at_from: str) -> int:
|
||||
return 19
|
||||
|
||||
def fake_post(*args, **kwargs):
|
||||
return MockResponse(200, {"text": "ok", "usage": {}})
|
||||
|
||||
monkeypatch.setattr("app.count_since", fake_count_since)
|
||||
monkeypatch.setattr(requests, "post", fake_post)
|
||||
|
||||
client = app.test_client()
|
||||
response = client.post(
|
||||
"/receipt",
|
||||
headers=auth_header(),
|
||||
data={
|
||||
"prompt": "p",
|
||||
"image": make_image(),
|
||||
},
|
||||
content_type="multipart/form-data",
|
||||
environ_base={"REMOTE_ADDR": "198.51.100.11"},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
def test_metrics_record_duration_and_status(monkeypatch) -> None:
|
||||
db_file = os.path.join(tempfile.gettempdir(), "recipe-service-tests-metrics-duration.sqlite3")
|
||||
if os.path.exists(db_file):
|
||||
os.remove(db_file)
|
||||
|
||||
monkeypatch.setenv("RECIPE_METRICS_DB", db_file)
|
||||
monkeypatch.setenv("RECIPE_API_TOKEN", "test-token")
|
||||
|
||||
from metrics import initialize
|
||||
|
||||
initialize()
|
||||
|
||||
def fake_post(*args, **kwargs):
|
||||
return MockResponse(500, {"error": {"message": "provider detail"}})
|
||||
|
||||
monkeypatch.setattr(requests, "post", fake_post)
|
||||
client = app.test_client()
|
||||
response = client.post(
|
||||
"/receipt",
|
||||
headers=auth_header(),
|
||||
data={
|
||||
"prompt": "p",
|
||||
"image": make_image(),
|
||||
},
|
||||
content_type="multipart/form-data",
|
||||
)
|
||||
|
||||
assert response.status_code == 502
|
||||
connection = sqlite3.connect(db_file)
|
||||
try:
|
||||
row = connection.execute(
|
||||
"SELECT status_code, duration_ms, error FROM requests ORDER BY rowid DESC LIMIT 1"
|
||||
).fetchone()
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
assert row is not None
|
||||
status_code, duration_ms, error = row
|
||||
assert status_code == 502
|
||||
assert duration_ms >= 0
|
||||
assert error == "upstream_error"
|
||||
Reference in New Issue
Block a user