diff --git a/CHANGELOG.md b/CHANGELOG.md index fc95e34..0658c4c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Added +- **Bounded retries and connection reuse for library-desk** - GETs and the read-only `POST /query/*` and `POST /rag/search` endpoints retry once (2 attempts, short backoff) on transport errors and retryable 5xx; wiki writes are never retried. The client now honors `LIBRARY_DESK_TIMEOUT` instead of hardcoded 60s/30s, a librarian run holds one shared HTTP connection instead of constructing a client per tool call, and read tools raise `ModelRetry` on transient HTTP errors so the agent's retry budget engages - **One librarian timeout budget** - new `LIBRARIAN_TIMEOUT` (default 180s) enforced with `asyncio.wait_for` inside `delegate_to_librarian`, capping the previously uncapped live paths (steward direct delegation and streaming). The Ollama provider's AsyncOpenAI client now carries an explicit `OLLAMA_TIMEOUT` instead of the SDK's ~600s default, and the contradictory unused 60s default in `AgentRequest.timeout_seconds` was removed (None defers to the configured budget) - **Search degradation signaling** - The librarian client parses `source_counts` (plus the additive `source_status`/`degraded` fields when a newer library-desk sends them; absence is tolerated), and `hybrid_search` appends a one-line coverage note when a search is degraded or an enabled source leg contributed nothing, so outages are visible to the model and the user diff --git a/src/agents/librarian/agent.py b/src/agents/librarian/agent.py index 3e57ba6..a926629 100644 --- a/src/agents/librarian/agent.py +++ b/src/agents/librarian/agent.py @@ -11,6 +11,7 @@ from typing import Any from pydantic_ai import Agent +from src.agents.librarian.client import library_client_session from src.agents.librarian.tools import ( create_wiki_page, explore_knowledge_graph, @@ -245,10 +246,12 @@ async def run_librarian( ) try: - result = await agent.run( - prompt, - message_history=message_history, - ) + # One shared library-desk connection for all tool calls in this run + async with library_client_session(): + result = await agent.run( + prompt, + message_history=message_history, + ) logger.info( "librarian_task_completed", @@ -311,12 +314,14 @@ async def run_librarian_stream( ) try: - async with agent.run_stream( - prompt, - message_history=message_history, - ) as response: - async for delta in response.stream_text(delta=True): - yield delta + # One shared library-desk connection for all tool calls in this run + async with library_client_session(): + async with agent.run_stream( + prompt, + message_history=message_history, + ) as response: + async for delta in response.stream_text(delta=True): + yield delta logger.info("librarian_stream_completed", task=task[:50]) diff --git a/src/agents/librarian/client.py b/src/agents/librarian/client.py index 809fd09..87d6f00 100644 --- a/src/agents/librarian/client.py +++ b/src/agents/librarian/client.py @@ -7,6 +7,10 @@ Provides async methods for all relevant library-desk endpoints: - Vector search - Knowledge graph queries """ +import asyncio +from collections.abc import AsyncIterator, Awaitable, Callable +from contextlib import asynccontextmanager +from contextvars import ContextVar from typing import Any import httpx @@ -18,6 +22,17 @@ from src.core.logging_config import get_logger logger = get_logger(__name__) +# Retry policy for idempotent/read-only requests (GETs, POST /query/*, +# POST /rag/search). Writes are never retried. +_RETRY_ATTEMPTS = 2 +_RETRY_BACKOFF_SECONDS = 0.5 +_RETRYABLE_STATUS_CODES = {502, 503, 504} + +# One shared HTTP connection per librarian run (see library_client_session) +_shared_http_client: ContextVar[httpx.AsyncClient | None] = ContextVar( + "library_desk_http_client", default=None +) + # ============================================================================ # Response Models @@ -181,7 +196,7 @@ class LibraryDeskClient: self, base_url: str | None = None, api_key: str | None = None, - timeout: int = 60, + timeout: int | None = None, ): """ Initialize the client. @@ -190,30 +205,55 @@ class LibraryDeskClient: base_url: Library-desk API URL (defaults to config) api_key: API key for authentication (defaults to config) timeout: Request timeout in seconds + (defaults to config.LIBRARY_DESK_TIMEOUT) """ self.base_url = base_url or str(config.LIBRARY_DESK_HOST) self.api_key = api_key or config.LIBRARY_DESK_API_KEY - self.timeout = timeout + self.timeout = timeout if timeout is not None else config.LIBRARY_DESK_TIMEOUT self._client: httpx.AsyncClient | None = None + self._owns_client = False - async def __aenter__(self) -> "LibraryDeskClient": - """Create HTTP client on context entry.""" + def _build_http_client(self) -> httpx.AsyncClient: + """Build a configured httpx client.""" headers = {} if self.api_key: headers["Authorization"] = f"Bearer {self.api_key}" - self._client = httpx.AsyncClient( + return httpx.AsyncClient( base_url=self.base_url, headers=headers, timeout=self.timeout, ) + + def _uses_default_target(self) -> bool: + """Whether this client targets the configured library-desk instance.""" + return ( + self.base_url == str(config.LIBRARY_DESK_HOST) + and self.api_key == config.LIBRARY_DESK_API_KEY + ) + + async def __aenter__(self) -> "LibraryDeskClient": + """ + Acquire an HTTP client on context entry. + + Reuses the run-level shared connection (see library_client_session) + when one is active, instead of constructing a new client per call. + """ + shared = _shared_http_client.get() + if shared is not None and not shared.is_closed and self._uses_default_target(): + self._client = shared + self._owns_client = False + else: + self._client = self._build_http_client() + self._owns_client = True return self async def __aexit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None: - """Close HTTP client on context exit.""" - if self._client: + """Close HTTP client on context exit (only if we own it).""" + if self._client and self._owns_client: await self._client.aclose() - self._client = None + self._client = None + self._owns_client = False def _ensure_client(self) -> httpx.AsyncClient: """Ensure client is initialized.""" @@ -223,6 +263,47 @@ class LibraryDeskClient: ) return self._client + async def _request_with_retry( + self, + send: Callable[[], Awaitable[httpx.Response]], + description: str, + ) -> httpx.Response: + """ + Send an idempotent/read-only request with a bounded retry. + + Retries once (2 attempts total) with a short backoff on transport + errors and retryable 5xx statuses. Only used for GETs and the + read-only POST /query/* and /rag/search endpoints - never for + wiki writes. + """ + for attempt in range(1, _RETRY_ATTEMPTS + 1): + try: + response = await send() + except httpx.TransportError as e: + if attempt >= _RETRY_ATTEMPTS: + raise + logger.warning( + "library_desk_retry", + request=description, + error=str(e), + attempt=attempt, + ) + else: + if ( + response.status_code not in _RETRYABLE_STATUS_CODES + or attempt >= _RETRY_ATTEMPTS + ): + return response + logger.warning( + "library_desk_retry", + request=description, + status_code=response.status_code, + attempt=attempt, + ) + await asyncio.sleep(_RETRY_BACKOFF_SECONDS * attempt) + + raise RuntimeError("unreachable") # pragma: no cover + # ======================================================================== # HybridRAG # ======================================================================== @@ -279,10 +360,13 @@ class LibraryDeskClient: logger.info("library_desk_hybrid_search", query=query, user=user) - response = await client.post( - "/query/hybrid", - json=payload, - params={"user": user}, + response = await self._request_with_retry( + lambda: client.post( + "/query/hybrid", + json=payload, + params={"user": user}, + ), + "POST /query/hybrid", ) response.raise_for_status() @@ -367,9 +451,12 @@ class LibraryDeskClient: logger.debug("library_desk_wiki_search", query=query, user=user) - response = await client.get( - "/wiki/search", - params={"q": query, "user": user, "limit": limit}, + response = await self._request_with_retry( + lambda: client.get( + "/wiki/search", + params={"q": query, "user": user, "limit": limit}, + ), + "GET /wiki/search", ) response.raise_for_status() @@ -394,9 +481,12 @@ class LibraryDeskClient: user = user or get_user() client = self._ensure_client() - response = await client.get( - f"/wiki/pages/{page_id}", - params={"user": user}, + response = await self._request_with_retry( + lambda: client.get( + f"/wiki/pages/{page_id}", + params={"user": user}, + ), + f"GET /wiki/pages/{page_id}", ) response.raise_for_status() @@ -426,7 +516,10 @@ class LibraryDeskClient: if tag: params["tag"] = tag - response = await client.get("/wiki/pages", params=params) + response = await self._request_with_retry( + lambda: client.get("/wiki/pages", params=params), + "GET /wiki/pages", + ) response.raise_for_status() data = response.json() @@ -612,9 +705,12 @@ class LibraryDeskClient: user = user or get_user() client = self._ensure_client() - response = await client.get( - "/wiki/dossiers", - params={"user": user}, + response = await self._request_with_retry( + lambda: client.get( + "/wiki/dossiers", + params={"user": user}, + ), + "GET /wiki/dossiers", ) response.raise_for_status() @@ -725,7 +821,10 @@ class LibraryDeskClient: if node_type: params["node_type"] = node_type - response = await client.get("/graph/nodes", params=params) + response = await self._request_with_retry( + lambda: client.get("/graph/nodes", params=params), + "GET /graph/nodes", + ) response.raise_for_status() data = response.json() @@ -749,9 +848,12 @@ class LibraryDeskClient: user = user or get_user() client = self._ensure_client() - response = await client.get( - f"/graph/nodes/{node_id}", - params={"user": user}, + response = await self._request_with_retry( + lambda: client.get( + f"/graph/nodes/{node_id}", + params={"user": user}, + ), + f"GET /graph/nodes/{node_id}", ) response.raise_for_status() @@ -770,7 +872,10 @@ class LibraryDeskClient: """ try: client = self._ensure_client() - response = await client.get("/health") + response = await self._request_with_retry( + lambda: client.get("/health"), + "GET /health", + ) return response.status_code == 200 except Exception as e: logger.warning("library_desk_health_check_failed", error=str(e)) @@ -814,7 +919,10 @@ class LibraryDeskClient: logger.info("library_desk_web_search", query=query, limit=limit) - response = await client.post("/rag/search", json=payload, timeout=30.0) + response = await self._request_with_retry( + lambda: client.post("/rag/search", json=payload), + "POST /rag/search", + ) response.raise_for_status() data = response.json() @@ -876,7 +984,7 @@ class LibraryDeskClient: logger.debug("library_desk_extract_content", url=url) - response = await client.post("/content/extract", json=payload, timeout=30.0) + response = await client.post("/content/extract", json=payload) response.raise_for_status() data = response.json() @@ -928,7 +1036,6 @@ class LibraryDeskClient: response = await client.post( "/content/extract/batch", json=payload, - timeout=60.0, # Longer timeout for batch ) response.raise_for_status() @@ -967,3 +1074,30 @@ async def get_library_client() -> LibraryDeskClient: results = await client.hybrid_search("query") """ return LibraryDeskClient() + + +@asynccontextmanager +async def library_client_session() -> AsyncIterator[None]: + """ + Hold ONE shared HTTP connection for the duration of a librarian run. + + While the session is active, every LibraryDeskClient targeting the + configured library-desk instance reuses the shared httpx client + instead of constructing (and tearing down) a connection per tool + call. Nested sessions are no-ops. + + Usage: + async with library_client_session(): + ... # librarian tools reuse one connection + """ + if _shared_http_client.get() is not None: + yield + return + + http_client = LibraryDeskClient()._build_http_client() + token = _shared_http_client.set(http_client) + try: + yield + finally: + _shared_http_client.reset(token) + await http_client.aclose() diff --git a/src/agents/librarian/tools.py b/src/agents/librarian/tools.py index b161983..095cb36 100644 --- a/src/agents/librarian/tools.py +++ b/src/agents/librarian/tools.py @@ -4,11 +4,33 @@ Librarian tools for PydanticAI agent. These tools wrap the library-desk API and are registered with The Librarian agent for research and knowledge management tasks. """ +import httpx +from pydantic_ai import ModelRetry + from src.agents.librarian.client import HybridRAGResponse, LibraryDeskClient from src.core.logging_config import get_logger logger = get_logger(__name__) + +def _retry_if_transient(e: Exception, what: str) -> None: + """ + Convert transient HTTP errors into ModelRetry so the agent's + retry budget (Agent(retries=2)) engages instead of the tool + swallowing the failure. + + Only read tools call this - writes are never retried to avoid + duplicate wiki pages. + """ + retryable = isinstance(e, httpx.TransportError) + if isinstance(e, httpx.HTTPStatusError): + status = e.response.status_code + retryable = status >= 500 or status == 429 + if retryable: + raise ModelRetry( + f"{what} is temporarily unavailable; please retry." + ) from e + # Icons keyed by the values library-desk emits in each result's `sources` # list (search legs) and `source_type` (result origin). SOURCE_ICONS = { @@ -179,6 +201,7 @@ async def hybrid_search( except Exception as e: logger.error("librarian_hybrid_search_error", error=str(e), query=query) + _retry_if_transient(e, "The knowledge archive") return "I was unable to search the knowledge archives; the search service did not respond properly." @@ -227,6 +250,7 @@ async def search_wiki( except Exception as e: logger.error("librarian_wiki_search_error", error=str(e)) + _retry_if_transient(e, "The wiki search") return "I was unable to search the wiki at this time." @@ -270,6 +294,7 @@ async def get_wiki_page( except Exception as e: logger.error("librarian_get_page_error", error=str(e), page_id=page_id) + _retry_if_transient(e, "The wiki") return f"I was unable to retrieve wiki page {page_id}." @@ -304,6 +329,7 @@ async def list_dossiers() -> str: except Exception as e: logger.error("librarian_list_dossiers_error", error=str(e)) + _retry_if_transient(e, "The dossier index") return "I was unable to retrieve the list of dossiers." @@ -345,6 +371,7 @@ async def get_dossier_pages( except Exception as e: logger.error("librarian_get_dossier_error", error=str(e)) + _retry_if_transient(e, "The dossier index") return f"I was unable to retrieve the dossier '{dossier_name}'." @@ -394,6 +421,7 @@ async def semantic_search( except Exception as e: logger.error("librarian_semantic_search_error", error=str(e)) + _retry_if_transient(e, "The semantic search") return "I was unable to complete the semantic search." @@ -447,6 +475,7 @@ async def explore_knowledge_graph( except Exception as e: logger.error("librarian_explore_graph_error", error=str(e)) + _retry_if_transient(e, "The knowledge graph") return "I was unable to explore the knowledge graph." @@ -518,6 +547,7 @@ async def find_related_entities( except Exception as e: logger.error("librarian_find_related_error", error=str(e)) + _retry_if_transient(e, "The knowledge graph") return f"I was unable to look up entities related to '{entity_name}'." @@ -601,6 +631,7 @@ async def search_web( except Exception as e: logger.error("librarian_web_search_error", error=str(e), query=query) + _retry_if_transient(e, "The web search") return "I was unable to search the web at this time." @@ -677,6 +708,7 @@ async def read_url( except Exception as e: logger.error("librarian_read_url_error", error=str(e), url=url) + _retry_if_transient(e, "Content extraction") return f"I was unable to read the page at {url}." @@ -751,6 +783,7 @@ async def read_urls_batch( except Exception as e: logger.error("librarian_read_urls_batch_error", error=str(e)) + _retry_if_transient(e, "Content extraction") return "I was unable to read the requested pages." diff --git a/tests/agents/librarian/test_client_retry.py b/tests/agents/librarian/test_client_retry.py new file mode 100644 index 0000000..23035fb --- /dev/null +++ b/tests/agents/librarian/test_client_retry.py @@ -0,0 +1,232 @@ +""" +Tests for bounded retries, timeout wiring, and client reuse in +LibraryDeskClient, plus ModelRetry escalation from the read tools. +""" + +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +from pydantic_ai import ModelRetry + +from src.agents.librarian.client import ( + LibraryDeskClient, + library_client_session, +) +from src.agents.librarian.tools import ( + create_wiki_page, + hybrid_search, + search_wiki, +) +from src.core.config import config + + +@pytest.fixture(autouse=True) +def _no_backoff(monkeypatch): + """Skip the retry backoff sleep in tests.""" + monkeypatch.setattr("src.agents.librarian.client._RETRY_BACKOFF_SECONDS", 0) + + +def _ok_response(payload: dict) -> MagicMock: + response = MagicMock() + response.status_code = 200 + response.json.return_value = payload + response.raise_for_status = MagicMock() + return response + + +@pytest.fixture +def client_with_mock(): + client = LibraryDeskClient(base_url="http://test:8089", api_key="test-key") + client._client = AsyncMock(spec=httpx.AsyncClient) + return client + + +@pytest.mark.unit +class TestBoundedRetries: + """2-attempt retry for GETs and read-only POST /query/*, /rag/search.""" + + @pytest.mark.asyncio + async def test_get_retries_once_on_transport_error(self, client_with_mock): + mock_httpx = client_with_mock._client + mock_httpx.get.side_effect = [ + httpx.ConnectError("Connection refused"), + _ok_response({"results": []}), + ] + + results = await client_with_mock.search_wiki("docker", user="u") + + assert results == [] + assert mock_httpx.get.call_count == 2 + + @pytest.mark.asyncio + async def test_get_gives_up_after_two_attempts(self, client_with_mock): + mock_httpx = client_with_mock._client + mock_httpx.get.side_effect = httpx.ConnectError("Connection refused") + + with pytest.raises(httpx.ConnectError): + await client_with_mock.search_wiki("docker", user="u") + + assert mock_httpx.get.call_count == 2 + + @pytest.mark.asyncio + async def test_query_hybrid_retries_on_503(self, client_with_mock): + mock_httpx = client_with_mock._client + bad = MagicMock() + bad.status_code = 503 + mock_httpx.post.side_effect = [ + bad, + _ok_response({"results": [], "keywords": {}, "context": ""}), + ] + + response = await client_with_mock.hybrid_search("docker", user="u") + + assert response.results == [] + assert mock_httpx.post.call_count == 2 + + @pytest.mark.asyncio + async def test_wiki_write_is_never_retried(self, client_with_mock): + """POST /wiki/pages must not retry - it could duplicate pages.""" + mock_httpx = client_with_mock._client + mock_httpx.post.side_effect = httpx.ConnectError("Connection refused") + + with pytest.raises(httpx.ConnectError): + await client_with_mock.create_wiki_page( + title="T", path="/t", content="c", user="u" + ) + + assert mock_httpx.post.call_count == 1 + + @pytest.mark.asyncio + async def test_smart_create_is_never_retried(self, client_with_mock): + mock_httpx = client_with_mock._client + mock_httpx.post.side_effect = httpx.ConnectError("Connection refused") + + with pytest.raises(httpx.ConnectError): + await client_with_mock.smart_create_wiki_page( + topic="T", tags=["x"], user="u" + ) + + assert mock_httpx.post.call_count == 1 + + +@pytest.mark.unit +class TestTimeoutWiring: + """LIBRARY_DESK_TIMEOUT config replaces the hardcoded 60s/30s.""" + + def test_default_timeout_from_config(self): + client = LibraryDeskClient() + + assert client.timeout == config.LIBRARY_DESK_TIMEOUT + + def test_explicit_timeout_wins(self): + client = LibraryDeskClient(timeout=5) + + assert client.timeout == 5 + + +@pytest.mark.unit +class TestClientReuse: + """One shared HTTP connection per librarian run.""" + + @pytest.mark.asyncio + async def test_clients_share_connection_inside_session(self): + async with library_client_session(): + async with LibraryDeskClient() as c1: + http1 = c1._client + # shared connection survives client exit + assert http1 is not None + assert not http1.is_closed + + async with LibraryDeskClient() as c2: + assert c2._client is http1 + + # session close tears the shared connection down + assert http1.is_closed + + @pytest.mark.asyncio + async def test_nested_sessions_are_noops(self): + async with library_client_session(): + async with LibraryDeskClient() as c1: + http1 = c1._client + async with library_client_session(): + async with LibraryDeskClient() as c2: + assert c2._client is http1 + # inner session exit must not close the shared connection + assert not http1.is_closed + + @pytest.mark.asyncio + async def test_client_owns_connection_outside_session(self): + async with LibraryDeskClient() as client: + http_client = client._client + + assert http_client.is_closed + + @pytest.mark.asyncio + async def test_custom_target_does_not_reuse_shared(self): + async with library_client_session(): + async with LibraryDeskClient() as shared_client: + shared_http = shared_client._client + async with LibraryDeskClient(base_url="http://other:9999") as custom: + assert custom._client is not shared_http + + +@pytest.mark.unit +class TestModelRetryEscalation: + """Read tools raise ModelRetry on transient errors so Agent(retries=2) engages.""" + + def _patched_client(self, mock_client): + factory = MagicMock() + factory.return_value.__aenter__ = AsyncMock(return_value=mock_client) + factory.return_value.__aexit__ = AsyncMock(return_value=None) + return patch("src.agents.librarian.tools.LibraryDeskClient", factory) + + @pytest.mark.asyncio + async def test_read_tool_raises_model_retry_on_transport_error(self): + mock_client = AsyncMock() + mock_client.hybrid_search.side_effect = httpx.ConnectError( + "Connection refused" + ) + + with self._patched_client(mock_client): + with pytest.raises(ModelRetry): + await hybrid_search("docker") + + @pytest.mark.asyncio + async def test_read_tool_raises_model_retry_on_5xx(self): + request = httpx.Request("GET", "http://test:8089/wiki/search") + response = httpx.Response(502, request=request) + mock_client = AsyncMock() + mock_client.search_wiki.side_effect = httpx.HTTPStatusError( + "bad gateway", request=request, response=response + ) + + with self._patched_client(mock_client): + with pytest.raises(ModelRetry): + await search_wiki("docker") + + @pytest.mark.asyncio + async def test_read_tool_returns_safe_message_on_non_transient(self): + mock_client = AsyncMock() + mock_client.hybrid_search.side_effect = ValueError("bad parse") + + with self._patched_client(mock_client): + result = await hybrid_search("docker") + + assert "unable" in result + assert "bad parse" not in result + + @pytest.mark.asyncio + async def test_write_tool_never_raises_model_retry(self): + mock_client = AsyncMock() + mock_client.create_wiki_page.side_effect = httpx.ConnectError( + "Connection refused" + ) + + with self._patched_client(mock_client): + result = await create_wiki_page( + title="T", path="/t", content="c", tags=["x"] + ) + + assert "unable" in result + assert "Connection refused" not in result