Skip to content
Open
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
13 changes: 11 additions & 2 deletions pageindex/page_index_md.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,8 +18,17 @@ async def get_node_summary(node, summary_token_threshold=200, model=None):

async def generate_summaries_for_structure_md(structure, summary_token_threshold, model=None):
nodes = structure_to_list(structure)
tasks = [get_node_summary(node, summary_token_threshold=summary_token_threshold, model=model) for node in nodes]
summaries = await asyncio.gather(*tasks)
tasks = [asyncio.create_task(get_node_summary(
node, summary_token_threshold=summary_token_threshold, model=model))
for node in nodes]
try:
summaries = await asyncio.gather(*tasks)
except BaseException:
# A failed gather leaves siblings running on a persistent event loop.
for task in tasks:
task.cancel()
await asyncio.gather(*tasks, return_exceptions=True)
raise

for node, summary in zip(nodes, summaries):
if not node.get('nodes'):
Expand Down
131 changes: 131 additions & 0 deletions tests/test_page_index_md.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,12 @@
import asyncio
import http.server
import json
import select
import socket
import tempfile
import threading
import time
from pathlib import Path
import unittest

from pageindex.page_index_md import extract_nodes_from_markdown
Expand All @@ -18,6 +27,128 @@ def test_skips_bold_heading_with_only_whitespace(self):
],
)

class _MarkdownProviderServer(http.server.ThreadingHTTPServer):
daemon_threads = True

def __init__(self, mode):
super().__init__(("127.0.0.1", 0), _MarkdownProviderHandler)
self.mode = mode
self.connections = []
self.received = threading.Event()
self.release = threading.Event()
self.lock = threading.Lock()

class _MarkdownProviderHandler(http.server.BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"

def log_message(self, *args):
pass

def do_POST(self):
request = json.loads(self.rfile.read(int(self.headers["Content-Length"])))
failing = "# Failure" in request["messages"][-1]["content"]
with self.server.lock:
self.server.connections.append(self.connection)
if len(self.server.connections) == 2:
self.server.received.set()
if not self.server.received.wait(5):
raise AssertionError("Summary requests did not overlap")
status = 200
if self.server.mode == "failure" and failing:
status = 404
body = {"error": {"message": "fixture rejected model", "type": "invalid_request_error", "code": "model_not_found"}}
else:
if self.server.mode != "success":
self.server.release.wait(10)
body = {"id": "fixture", "object": "chat.completion", "created": 0, "model": "gpt-4o-mini", "choices": [{"index": 0, "message": {"role": "assistant", "content": "Summary"}, "finish_reason": "stop"}], "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}}
data = json.dumps(body).encode()
try:
self.send_response(status)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(data)))
self.send_header("Connection", "close")
self.end_headers()
self.wfile.write(data)
except (BrokenPipeError, ConnectionResetError):
pass

async def _wait_markdown_provider(predicate, timeout=3):
deadline = time.monotonic() + timeout
while not predicate():
if time.monotonic() >= deadline:
return False
await asyncio.sleep(0.01)
return True

def _markdown_provider_peer_open(connection):
if connection.fileno() < 0:
return False
try:
readable, _, _ = select.select([connection], [], [], 0)
return not readable or connection.recv(1, socket.MSG_PEEK) != b""
except OSError:
return False

async def _native_markdown_summaries(mode):
from pageindex import md_to_tree
from pageindex.utils import _llm_backend
server = _MarkdownProviderServer(mode)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
before = asyncio.all_tasks()
token = _llm_backend.set({"api_base": f"http://127.0.0.1:{server.server_port}/v1", "api_key": "fixture-only"})
try:
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "note.md"
path.write_text("# Failure\n\nFirst body.\n\n# Slow\n\nSecond body.\n")
task = asyncio.create_task(md_to_tree(path, if_add_node_summary="yes", summary_token_threshold=0, model="gpt-4o-mini"))
assert await _wait_markdown_provider(server.received.is_set, 8), "Actual provider HTTP requests must run concurrently"
if mode == "cancel":
task.cancel()
error = None
try:
result = await asyncio.wait_for(task, 0.05) if mode == "timeout" else await task
except BaseException as exc:
error = type(exc).__name__

await _wait_markdown_provider(lambda: not any(_markdown_provider_peer_open(c) for c in server.connections), 0.3)
pending = [t for t in asyncio.all_tasks() - before if not t.done() and "get_node_summary" in t.get_coro().__qualname__]
out = {"mode": mode, "requests": len(server.connections), "public_error": error, "pending_owned_summary_tasks": len(pending), "open_provider_peers": sum(_markdown_provider_peer_open(c) for c in server.connections)}
if mode == "success":
assert all(node["summary"] == "Summary" for node in result["structure"])
# The proof owns cleanup even on unchanged source; no request is left alive.
for own in pending:
own.cancel()
await asyncio.gather(*pending, return_exceptions=True)
return out
finally:
_llm_backend.reset(token)
server.release.set()
server.shutdown()
server.server_close()
thread.join()


class MarkdownSummaryTaskCleanupTest(unittest.TestCase):
def check_cleanup(self, mode, expected_error):
result = asyncio.run(_native_markdown_summaries(mode))
self.assertEqual(result["requests"], 2) # both native requests overlapped
self.assertEqual(result["public_error"], expected_error)
self.assertEqual(result["pending_owned_summary_tasks"], 0)
self.assertEqual(result["open_provider_peers"], 0)

def test_success_preserves_concurrent_summaries(self):
self.check_cleanup("success", None)

def test_fatal_provider_error_closes_sibling_request_before_return(self):
self.check_cleanup("failure", "NotFoundError")

def test_caller_cancellation_closes_owned_requests(self):
self.check_cleanup("cancel", "CancelledError")

def test_caller_timeout_closes_owned_requests(self):
self.check_cleanup("timeout", "TimeoutError")


class MarkdownCliTest(unittest.TestCase):
def test_md_cli_runs_without_llm_or_key(self):
Expand Down