import asyncio
import http.server
import threading
import time
import pytest
from conftest import _ThreadingHTTPServer
import eggfetch
class _SlowHandler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
if self.path.startswith("/slow"):
time.sleep(0.5) body = b"OK"
self.send_response(200)
self.send_header("Content-Type", "text/plain")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def log_message(self, format, *args):
pass
@pytest.fixture(scope="module")
def slow_server():
srv = _ThreadingHTTPServer(("127.0.0.1", 0), _SlowHandler)
port = srv.server_address[1]
t = threading.Thread(target=srv.serve_forever, daemon=True)
t.start()
yield f"http://127.0.0.1:{port}"
srv.shutdown()
class TestSyncCloseDuringRequest:
def test_close_during_concurrent_requests(self, slow_server):
client = eggfetch.Client()
errors = []
def make_request(idx):
try:
r = client.get(f"{slow_server}/slow")
except ValueError as e:
if "closed" not in str(e):
errors.append(f"thread {idx}: unexpected error: {e}")
except Exception:
pass
threads = [threading.Thread(target=make_request, args=(i,)) for i in range(5)]
for t in threads:
t.start()
time.sleep(0.05) client.close()
for t in threads:
t.join(timeout=5)
assert not errors, f"errors occurred: {errors}"
def test_close_is_idempotent_under_contention(self, slow_server):
client = eggfetch.Client()
errors = []
def close_client(idx):
try:
client.close()
except Exception as e:
errors.append(f"thread {idx}: unexpected error: {e}")
threads = [threading.Thread(target=close_client, args=(i,)) for i in range(10)]
for t in threads:
t.start()
for t in threads:
t.join(timeout=5)
assert client.is_closed
assert not errors, f"errors occurred: {errors}"
def test_request_after_close_from_other_thread(self, slow_server):
client = eggfetch.Client()
results = []
close_event = threading.Event()
def make_request():
try:
r = client.get(f"{slow_server}/slow")
results.append("success")
except ValueError as e:
if "closed" in str(e):
results.append("closed")
else:
results.append(f"error: {e}")
except Exception:
pass
def close_later():
close_event.wait()
client.close()
req_thread = threading.Thread(target=make_request)
close_thread = threading.Thread(target=close_later)
req_thread.start()
close_thread.start()
close_event.set() req_thread.join(timeout=5)
close_thread.join(timeout=5)
assert len(results) <= 1
if results:
assert results[0] in ("success", "closed")
class TestAsyncCloseDuringRequest:
def test_close_during_concurrent_requests(self, slow_server):
async def _test():
client = eggfetch.AsyncClient()
tasks = [client.get(f"{slow_server}/slow") for _ in range(5)]
running = [asyncio.ensure_future(t) for t in tasks]
await asyncio.sleep(0.05)
client.close()
results = await asyncio.gather(*running, return_exceptions=True)
for r in results:
if isinstance(r, Exception):
assert "closed" in str(r) or "connect" in str(r).lower()
asyncio.run(_test())
def test_close_is_idempotent_under_contention(self, slow_server):
async def _test():
client = eggfetch.AsyncClient()
async def close_client():
client.close()
await asyncio.gather(*[close_client() for _ in range(10)])
assert client.is_closed
asyncio.run(_test())
def test_request_after_close_from_other_task(self, slow_server):
async def _test():
client = eggfetch.AsyncClient()
results = []
async def make_request():
try:
r = await client.get(f"{slow_server}/slow")
results.append("success")
except ValueError as e:
if "closed" in str(e):
results.append("closed")
else:
results.append(f"error: {e}")
except Exception as e:
results.append(f"error: {e}")
async def close_later():
await asyncio.sleep(0.05)
client.close()
req_task = asyncio.ensure_future(make_request())
close_task = asyncio.ensure_future(close_later())
await asyncio.gather(req_task, close_task, return_exceptions=True)
assert len(results) == 1
assert results[0] in ("success", "closed")
asyncio.run(_test())
class TestContextManagerCloseDuringRequest:
def test_sync_context_exit_during_request(self, slow_server):
errors = []
def make_request(client):
try:
r = client.get(f"{slow_server}/slow")
except ValueError as e:
if "closed" not in str(e):
errors.append(f"unexpected error: {e}")
except Exception:
pass
with eggfetch.Client() as client:
threads = [threading.Thread(target=make_request, args=(client,)) for _ in range(3)]
for t in threads:
t.start()
time.sleep(0.05)
for t in threads:
t.join(timeout=5)
assert not errors, f"errors occurred: {errors}"
def test_async_context_exit_during_request(self, slow_server):
async def _test():
errors = []
async with eggfetch.AsyncClient() as client:
tasks = [client.get(f"{slow_server}/slow") for _ in range(3)]
running = [asyncio.ensure_future(t) for t in tasks]
await asyncio.sleep(0.05)
results = await asyncio.gather(*running, return_exceptions=True)
for r in results:
if isinstance(r, Exception):
if "closed" not in str(r) and "connect" not in str(r).lower():
errors.append(f"unexpected error: {r}")
assert not errors, f"errors occurred: {errors}"
asyncio.run(_test())