mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-25 14:06:18 +00:00
fix(persistence): drain schema connection close across cancellation (#5617)
This commit is contained in:
parent
1e3bfa09d4
commit
45cd0450b9
@ -22,6 +22,8 @@ from __future__ import annotations
|
||||
import re
|
||||
from urllib.parse import parse_qsl, quote, urlencode, urlsplit, urlunsplit
|
||||
|
||||
from deerflow.utils.file_io import await_drained
|
||||
|
||||
|
||||
def build_asyncpg_connect_args(schema: str) -> dict:
|
||||
"""Return SQLAlchemy ``connect_args`` that pin asyncpg's search_path.
|
||||
@ -254,4 +256,4 @@ async def ensure_postgres_schema_async(conn_string: str, schema: str, *, install
|
||||
try:
|
||||
await conn.execute(statement)
|
||||
finally:
|
||||
await conn.close()
|
||||
await await_drained(conn.close())
|
||||
|
||||
@ -1,5 +1,8 @@
|
||||
"""Tests for the PostgreSQL schema helpers (Issue #3380)."""
|
||||
|
||||
import asyncio
|
||||
import sys
|
||||
from types import SimpleNamespace
|
||||
from urllib.parse import parse_qs, urlsplit
|
||||
|
||||
import pytest
|
||||
@ -9,6 +12,7 @@ from deerflow.persistence.postgres_schema import (
|
||||
build_psycopg_options,
|
||||
create_schema_sql,
|
||||
dsn_with_search_path,
|
||||
ensure_postgres_schema_async,
|
||||
normalize_libpq_dsn,
|
||||
)
|
||||
|
||||
@ -176,3 +180,64 @@ class TestNormalizeLibpqDsn:
|
||||
def test_rejects_non_postgres_scheme(self):
|
||||
with pytest.raises(ValueError, match="Unsupported PostgreSQL DSN scheme"):
|
||||
normalize_libpq_dsn("mysql://localhost/db")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_schema_close_drains_across_repeated_cancellation(monkeypatch) -> None:
|
||||
class _BlockingConnection:
|
||||
def __init__(self) -> None:
|
||||
self.execute_started = asyncio.Event()
|
||||
self.allow_execute = asyncio.Event()
|
||||
self.close_started = asyncio.Event()
|
||||
self.allow_close = asyncio.Event()
|
||||
self.close_finished = asyncio.Event()
|
||||
|
||||
async def execute(self, _statement: str) -> None:
|
||||
self.execute_started.set()
|
||||
await self.allow_execute.wait()
|
||||
|
||||
async def close(self) -> None:
|
||||
self.close_started.set()
|
||||
await self.allow_close.wait()
|
||||
self.close_finished.set()
|
||||
|
||||
conn = _BlockingConnection()
|
||||
|
||||
class _AsyncConnection:
|
||||
@staticmethod
|
||||
async def connect(_dsn: str, *, autocommit: bool):
|
||||
assert autocommit is True
|
||||
return conn
|
||||
|
||||
monkeypatch.setitem(sys.modules, "psycopg", SimpleNamespace(AsyncConnection=_AsyncConnection))
|
||||
task: asyncio.Task[None] | None = None
|
||||
|
||||
try:
|
||||
task = asyncio.create_task(
|
||||
ensure_postgres_schema_async(
|
||||
"postgresql://user:pass@localhost/deerflow",
|
||||
"deerflow",
|
||||
install_hint="install postgres extras",
|
||||
)
|
||||
)
|
||||
await asyncio.wait_for(conn.execute_started.wait(), timeout=1)
|
||||
|
||||
task.cancel()
|
||||
await asyncio.wait_for(conn.close_started.wait(), timeout=1)
|
||||
|
||||
task.cancel()
|
||||
for _ in range(5):
|
||||
await asyncio.sleep(0)
|
||||
assert not task.done(), "schema setup returned before psycopg connection close finished"
|
||||
assert not conn.close_finished.is_set()
|
||||
|
||||
conn.allow_close.set()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
assert conn.close_finished.is_set()
|
||||
finally:
|
||||
conn.allow_execute.set()
|
||||
conn.allow_close.set()
|
||||
if task is not None and not task.done():
|
||||
task.cancel()
|
||||
await asyncio.gather(task, return_exceptions=True)
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user