diff --git a/simpleaudit/_event_loop.py b/simpleaudit/_event_loop.py new file mode 100644 index 0000000..cfb2109 --- /dev/null +++ b/simpleaudit/_event_loop.py @@ -0,0 +1,40 @@ +"""Running SimpleAudit's coroutines from synchronous code.""" + +import asyncio +from typing import Any, Coroutine, Dict + + +def _is_stale_client_close(context: Dict[str, Any]) -> bool: + """Is this an HTTP client from an earlier run, closed on a loop that is gone? + + Each sync ``run()`` gets its own event loop. The OpenAI SDK's async client sits + in a reference cycle, so one from an earlier run can be garbage-collected while + a later run's loop is running. Its ``__del__`` then schedules ``aclose()`` on + that loop, for connections that were opened on the earlier, closed one, and the + task fails with "Event loop is closed". Nothing is lost: those connections ended + with their loop. But asyncio prints the traceback as "Task exception was never + retrieved", which reads like a failed audit. + + any-llm has no public way to close a provider's client, so the client cannot be + closed before its loop ends; until it does, the message is dropped here. + """ + exception = context.get("exception") + if not isinstance(exception, RuntimeError) or str(exception) != "Event loop is closed": + return False + task = context.get("future") or context.get("task") + get_coro = getattr(task, "get_coro", None) + coro = get_coro() if get_coro is not None else None + return getattr(coro, "__qualname__", "").endswith("AsyncClient.aclose") + + +def _exception_handler(loop: asyncio.AbstractEventLoop, context: Dict[str, Any]) -> None: + if _is_stale_client_close(context): + return + loop.default_exception_handler(context) + + +def run_sync(coro: Coroutine[Any, Any, Any]) -> Any: + """``asyncio.run(coro)``, minus the stale-client message described above.""" + with asyncio.Runner() as runner: + runner.get_loop().set_exception_handler(_exception_handler) + return runner.run(coro) diff --git a/simpleaudit/cross_judge.py b/simpleaudit/cross_judge.py index 700f3a1..f050dc9 100644 --- a/simpleaudit/cross_judge.py +++ b/simpleaudit/cross_judge.py @@ -20,6 +20,7 @@ from pathlib import Path from typing import Any, Dict, List, Optional, Union +from simpleaudit._event_loop import run_sync from simpleaudit.experiment import AuditExperiment from simpleaudit.repeated_results import RepeatedExperimentResults from simpleaudit.utils import SEVERITY_ORDER, severity_direction @@ -354,7 +355,7 @@ def run( try: asyncio.get_running_loop() except RuntimeError: - return asyncio.run( + return run_sync( self.run_async( scenarios=scenarios, max_turns=max_turns, diff --git a/simpleaudit/experiment.py b/simpleaudit/experiment.py index 8addeb7..6ee87dd 100644 --- a/simpleaudit/experiment.py +++ b/simpleaudit/experiment.py @@ -8,6 +8,7 @@ from collections import Counter from tqdm.auto import tqdm +from simpleaudit._event_loop import run_sync from simpleaudit.results import AuditResult, AuditResults from simpleaudit.model_auditor import ModelAuditor from simpleaudit.repeated_results import RepeatedExperimentResults @@ -689,7 +690,7 @@ def run( try: asyncio.get_running_loop() except RuntimeError: - return asyncio.run( + return run_sync( self.run_async( scenarios, max_turns=max_turns, diff --git a/simpleaudit/model_auditor.py b/simpleaudit/model_auditor.py index 4218ec9..041eb81 100644 --- a/simpleaudit/model_auditor.py +++ b/simpleaudit/model_auditor.py @@ -25,6 +25,7 @@ from any_llm import AnyLLM from tqdm.auto import tqdm +from ._event_loop import run_sync from .context_marks import render_documents from .judges import get_judge from .judges.compose import SEVERITY_RESPONSE_SCHEMA @@ -1222,7 +1223,7 @@ def run( try: asyncio.get_running_loop() except RuntimeError: - return asyncio.run( + return run_sync( self.run_async( scenarios, max_turns=max_turns, diff --git a/simpleaudit/reframing.py b/simpleaudit/reframing.py index 62bf703..92f8799 100644 --- a/simpleaudit/reframing.py +++ b/simpleaudit/reframing.py @@ -75,6 +75,7 @@ from pathlib import Path from typing import Any, Callable, Dict, List, Mapping, Optional, Sequence, Union +from simpleaudit._event_loop import run_sync from simpleaudit.judges import get_judge from simpleaudit.model_auditor import ModelAuditor from simpleaudit.repeated_results import ( @@ -1005,7 +1006,7 @@ def _run_sync(coro_factory: Callable[[], Any], name: str) -> Any: try: asyncio.get_running_loop() except RuntimeError: - return asyncio.run(coro_factory()) + return run_sync(coro_factory()) raise RuntimeError( f"{name}() cannot be called from an active event loop. " f"Use await {name}_async() instead." diff --git a/tests/test_event_loop.py b/tests/test_event_loop.py new file mode 100644 index 0000000..71a1c9a --- /dev/null +++ b/tests/test_event_loop.py @@ -0,0 +1,136 @@ +"""run_sync: one event loop per sync run, without the stale-client close message.""" + +import asyncio +import gc +import json +import logging +import threading +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer + +import pytest + +from simpleaudit._event_loop import _exception_handler, _is_stale_client_close, run_sync + +NEVER_RETRIEVED = "Task exception was never retrieved" + + +class AsyncClient: + """Stands in for httpx's client: only the coroutine's qualified name is checked.""" + + async def aclose(self): + raise RuntimeError("Event loop is closed") + + +class Other: + async def aclose(self): + raise RuntimeError("Event loop is closed") + + +def _context(coro, exception): + loop = asyncio.new_event_loop() + try: + task = loop.create_task(coro) + loop.run_until_complete(asyncio.wait([task])) + return {"message": NEVER_RETRIEVED, "exception": exception, "future": task} + finally: + loop.close() + + +class TestFilter: + def test_a_client_close_on_a_closed_loop_is_dropped(self): + ctx = _context(AsyncClient().aclose(), RuntimeError("Event loop is closed")) + assert _is_stale_client_close(ctx) + + def test_another_coroutine_failing_the_same_way_is_kept(self): + ctx = _context(Other().aclose(), RuntimeError("Event loop is closed")) + assert not _is_stale_client_close(ctx) + + def test_a_client_close_failing_differently_is_kept(self): + ctx = _context(AsyncClient().aclose(), RuntimeError("connection reset")) + assert not _is_stale_client_close(ctx) + + def test_kept_contexts_reach_the_default_handler(self): + seen = [] + + class Loop: + def default_exception_handler(self, context): + seen.append(context) + + kept = {"message": "something else", "exception": ValueError("x")} + _exception_handler(Loop(), kept) + _exception_handler(Loop(), _context(AsyncClient().aclose(), RuntimeError("Event loop is closed"))) + assert seen == [kept] + + +def test_run_sync_returns_the_result_and_raises_the_error(): + async def ok(): + return 42 + + async def bad(): + raise ValueError("boom") + + assert run_sync(ok()) == 42 + with pytest.raises(ValueError, match="boom"): + run_sync(bad()) + + +class _Completions(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" # keep-alive, so the connection stays pooled + + def do_POST(self): + self.rfile.read(int(self.headers.get("Content-Length", 0))) + body = json.dumps({ + "id": "x", "object": "chat.completion", "created": 0, "model": "m", + "choices": [{"index": 0, "finish_reason": "stop", + "message": {"role": "assistant", "content": "hi"}}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + }).encode() + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def log_message(self, *args): + pass + + +@pytest.fixture +def base_url(): + server = ThreadingHTTPServer(("127.0.0.1", 0), _Completions) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + yield f"http://127.0.0.1:{server.server_address[1]}/v1" + server.shutdown() + server.server_close() + + +def _two_runs(runner, base_url): + """A client used in one run and garbage-collected during the next.""" + openai = pytest.importorskip("openai") + holder = {} + + async def first(): + client = openai.AsyncOpenAI(base_url=base_url, api_key="x", max_retries=0) + await client.chat.completions.create(model="m", messages=[{"role": "user", "content": "hi"}]) + holder["client"] = client + + async def second(): + holder.clear() + gc.collect() # the client sits in a reference cycle; its __del__ runs here + await asyncio.sleep(0.05) # the aclose() task it scheduled runs and fails + gc.collect() + + runner(first()) + runner(second()) + + +def test_a_client_from_an_earlier_run_closes_quietly(base_url, caplog): + caplog.set_level(logging.ERROR, logger="asyncio") + _two_runs(asyncio.run, base_url) + if NEVER_RETRIEVED not in caplog.text: + pytest.skip("this openai version no longer schedules aclose() from __del__") + caplog.clear() + + _two_runs(run_sync, base_url) + assert NEVER_RETRIEVED not in caplog.text