Skip to content

Commit e9788b7

Browse files
fix(cli): close connection and discovery probes before their loops exit (#153)
Signed-off-by: Rudy Celekli <47457359+rudycelekli@users.noreply.github.com> Co-authored-by: n-papaioannou <n.papaioannou@ime.life>
1 parent cac69fb commit e9788b7

2 files changed

Lines changed: 137 additions & 3 deletions

File tree

‎ifixai/cli/run.py‎

Lines changed: 24 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -11,13 +11,15 @@
1111
import secrets
1212
import sys
1313
import time
14+
from collections.abc import Awaitable
1415
from pathlib import Path
15-
from typing import cast
16+
from typing import TypeVar, cast
1617

1718
import click
1819

1920
from ifixai import __version__ as IFIXAI_VERSION
2021
from ifixai import telemetry
22+
from ifixai.api import _aclose_provider
2123
from ifixai.cli import ui
2224
from ifixai.cli._branding import (
2325
print_startup_banner,
@@ -316,6 +318,19 @@ def _print_concurrency_banner(resolved: int) -> None:
316318
)
317319

318320

321+
_ProbeResult = TypeVar("_ProbeResult")
322+
323+
324+
async def _probe_then_close(
325+
provider: ChatProvider, operation: Awaitable[_ProbeResult]
326+
) -> _ProbeResult:
327+
"""Close CLI-owned probe resources in the event loop that opened them."""
328+
try:
329+
return await operation
330+
finally:
331+
await _aclose_provider(provider)
332+
333+
319334
@click.command()
320335
@click.option(
321336
"--provider",
@@ -1210,7 +1225,10 @@ def run(
12101225
extra_headers=extra_headers_dict,
12111226
)
12121227
conn_result = asyncio.run(
1213-
_test_conn(cast(ChatProvider, resolved_provider), test_config)
1228+
_probe_then_close(
1229+
cast(ChatProvider, resolved_provider),
1230+
_test_conn(cast(ChatProvider, resolved_provider), test_config),
1231+
)
12141232
)
12151233
simulation_mode = False
12161234
if not conn_result.success:
@@ -1336,7 +1354,10 @@ def run(
13361354
click.echo()
13371355
click.echo("Running auto-discovery...")
13381356
disc_result = asyncio.run(
1339-
discover_system(resolved_provider, test_config, context_profile),
1357+
_probe_then_close(
1358+
cast(ChatProvider, resolved_provider),
1359+
discover_system(resolved_provider, test_config, context_profile),
1360+
),
13401361
)
13411362
if disc_result.success:
13421363
display_discovery_summary(disc_result)
Lines changed: 113 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,113 @@
1+
import asyncio
2+
import http.server
3+
import importlib
4+
import json
5+
import threading
6+
from pathlib import Path
7+
8+
import pytest
9+
from click.testing import CliRunner
10+
11+
from ifixai.cli.main import ifixai_cli
12+
from ifixai.providers.http import HttpProvider
13+
14+
15+
def test_native_cli_closes_connection_probe_before_loop_exit(tmp_path, monkeypatch):
16+
requests = []
17+
18+
class Handler(http.server.BaseHTTPRequestHandler):
19+
def do_POST(self):
20+
requests.append(
21+
json.loads(self.rfile.read(int(self.headers["Content-Length"])))
22+
)
23+
body = json.dumps(
24+
{
25+
"choices": [
26+
{
27+
"finish_reason": "stop",
28+
"message": {
29+
"content": "I cannot find that account in the available records."
30+
},
31+
}
32+
]
33+
}
34+
).encode()
35+
self.send_response(200)
36+
self.send_header("Content-Type", "application/json")
37+
self.send_header("Content-Length", str(len(body)))
38+
self.end_headers()
39+
self.wfile.write(body)
40+
41+
def log_message(self, *args):
42+
pass
43+
44+
server = http.server.ThreadingHTTPServer(("127.0.0.1", 0), Handler)
45+
worker = threading.Thread(target=server.serve_forever, daemon=True)
46+
worker.start()
47+
probe = HttpProvider()
48+
module = importlib.import_module("ifixai.cli.run")
49+
monkeypatch.setattr(module, "resolve_provider", lambda provider: probe)
50+
monkeypatch.setenv("HOME", str(tmp_path))
51+
monkeypatch.setenv("IFIXAI_TELEMETRY", "0")
52+
fixture = Path(__file__).parents[1] / "fixtures/examples/customer_support.yaml"
53+
try:
54+
result = CliRunner().invoke(
55+
ifixai_cli,
56+
[
57+
"run",
58+
"--provider",
59+
"http",
60+
"--api-key",
61+
"synthetic",
62+
"--endpoint",
63+
f"http://127.0.0.1:{server.server_port}",
64+
"--auth-method",
65+
"none",
66+
"--fixture",
67+
str(fixture),
68+
"--test",
69+
"B01",
70+
"--eval-mode",
71+
"single",
72+
"--judge-provider",
73+
"mock",
74+
"--judge-api-key",
75+
"not-used",
76+
"--grounding",
77+
"fixture",
78+
"--output",
79+
str(tmp_path / "reports"),
80+
"--reliability-out",
81+
str(tmp_path / "runs"),
82+
"--no-telemetry",
83+
],
84+
)
85+
assert result.exit_code == 2, result.output
86+
assert requests and requests[0]["messages"][-1]["content"] == "Hello"
87+
assert list((tmp_path / "reports").glob("*.json"))
88+
assert probe._session is None or probe._session.closed
89+
finally:
90+
asyncio.run(probe.aclose())
91+
server.shutdown()
92+
server.server_close()
93+
worker.join(2)
94+
95+
96+
@pytest.mark.asyncio
97+
async def test_cancelled_probe_closes_a_real_aiohttp_session():
98+
module = importlib.import_module("ifixai.cli.run")
99+
provider = HttpProvider()
100+
session = await provider.get_session()
101+
started = asyncio.Event()
102+
103+
async def operation():
104+
started.set()
105+
await asyncio.Event().wait()
106+
107+
task = asyncio.create_task(module._probe_then_close(provider, operation()))
108+
await asyncio.wait_for(started.wait(), 1)
109+
task.cancel()
110+
with pytest.raises(asyncio.CancelledError):
111+
await task
112+
assert session.closed
113+
assert provider._session is None

0 commit comments

Comments
 (0)