Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 40 additions & 0 deletions simpleaudit/_event_loop.py
Original file line number Diff line number Diff line change
@@ -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)
3 changes: 2 additions & 1 deletion simpleaudit/cross_judge.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down
3 changes: 2 additions & 1 deletion simpleaudit/experiment.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down
3 changes: 2 additions & 1 deletion simpleaudit/model_auditor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down
3 changes: 2 additions & 1 deletion simpleaudit/reframing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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."
Expand Down
136 changes: 136 additions & 0 deletions tests/test_event_loop.py
Original file line number Diff line number Diff line change
@@ -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
Loading