From 5bb0edbd012e1acd72af33c81b89add97ba87039 Mon Sep 17 00:00:00 2001 From: Joris Hartog Date: Mon, 29 Aug 2022 22:03:04 +0200 Subject: [PATCH] Add endpoints to cost/hit tests --- tests/test_fastapi_extension.py | 18 +++++++++++++++++ tests/test_starlette_extension.py | 32 +++++++++++++++++++++++++++++++ 2 files changed, 50 insertions(+) diff --git a/tests/test_fastapi_extension.py b/tests/test_fastapi_extension.py index 32225c0..4f2b666 100644 --- a/tests/test_fastapi_extension.py +++ b/tests/test_fastapi_extension.py @@ -300,11 +300,19 @@ class TestDecorators(TestSlowapi): async def t1(request: Request): return PlainTextResponse("test") + @app.get("/t2") + @limiter.limit("50/minute", cost=15) + async def t2(request: Request): + return PlainTextResponse("test") + client = TestClient(app) for i in range(0, 10): response = client.get("/t1") assert response.status_code == 200 if i < 5 else 429 + response = client.get("/t2") + assert response.status_code == 200 if i < 3 else 429 + def test_callable_cost(self): app, limiter = self.build_fastapi_app(key_func=get_ipaddr) @@ -313,7 +321,17 @@ class TestDecorators(TestSlowapi): async def t1(request: Request): return PlainTextResponse("test") + @app.get("/t2") + @limiter.limit( + "50/minute", cost=lambda request: int(request.headers["foo"]) * 1.5 + ) + async def t2(request: Request): + return PlainTextResponse("test") + client = TestClient(app) for i in range(0, 10): response = client.get("/t1", headers={"foo": "10"}) assert response.status_code == 200 if i < 5 else 429 + + response = client.get("/t2", headers={"foo": "5"}) + assert response.status_code == 200 if i < 6 else 429 diff --git a/tests/test_starlette_extension.py b/tests/test_starlette_extension.py index c02eb48..39058ad 100644 --- a/tests/test_starlette_extension.py +++ b/tests/test_starlette_extension.py @@ -269,12 +269,27 @@ class TestDecorators(TestSlowapi): app.add_route("/t1", t1) + @limiter.limit("50/minute", cost=15) + async def t2(request: Request): + return PlainTextResponse("test") + + app.add_route("/t2", t2) + client = TestClient(app) for i in range(0, 10): response = client.get("/t1") assert response.status_code == 200 if i < 5 else 429 if i < 5: assert response.text == "test" + else: + assert "error" in response.json() + + response = client.get("/t2") + assert response.status_code == 200 if i < 3 else 429 + if i < 3: + assert response.text == "test" + else: + assert "error" in response.json() def test_callable_cost(self): app, limiter = self.build_starlette_app(key_func=get_ipaddr) @@ -285,9 +300,26 @@ class TestDecorators(TestSlowapi): app.add_route("/t1", t1) + @limiter.limit( + "50/minute", cost=lambda request: int(request.headers["foo"]) * 1.5 + ) + async def t2(request: Request): + return PlainTextResponse("test") + + app.add_route("/t2", t2) + client = TestClient(app) for i in range(0, 10): response = client.get("/t1", headers={"foo": "10"}) assert response.status_code == 200 if i < 5 else 429 if i < 5: assert response.text == "test" + else: + assert "error" in response.json() + + response = client.get("/t2", headers={"foo": "5"}) + assert response.status_code == 200 if i < 6 else 429 + if i < 6: + assert response.text == "test" + else: + assert "error" in response.json()