feat(librarian): bounded retries, timeout wiring, and client reuse

- 2-attempt short-backoff retry for GETs and the read-only
  POST /query/* and /rag/search endpoints only; wiki writes are never
  retried (duplicate-page risk)
- honor the defined-but-ignored LIBRARY_DESK_TIMEOUT config instead of
  hardcoded 60s/30s per-call values
- hold ONE shared httpx.AsyncClient per librarian run via
  library_client_session (contextvar), instead of constructing a
  client per tool call; nested sessions are no-ops and custom targets
  still get their own client
- read tools raise ModelRetry on transient HTTP errors (transport
  errors, 5xx, 429) so Agent(retries=2) engages; write tools keep
  returning safe failure messages

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
2026-07-14 10:24:42 +02:00
co-authored by Claude Fable 5
parent 18f2e0efbd
commit 24fed8814f
5 changed files with 445 additions and 40 deletions
+15 -10
View File
@@ -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])
+164 -30
View File
@@ -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()
+33
View File
@@ -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."