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()