Use dict-keyed counter for concurrent-safe search cap
This commit is contained in:
parent
e2749ad2a6
commit
bf8eec68e3
1 changed files with 6 additions and 10 deletions
|
|
@ -32,10 +32,9 @@ def create_search_toolset(
|
||||||
Returns:
|
Returns:
|
||||||
FunctionToolset with a search tool.
|
FunctionToolset with a search tool.
|
||||||
"""
|
"""
|
||||||
# Per-run search counter. Resets when run_id changes so the toolset
|
# Per-run search counter keyed by run_id. Safe for concurrent runs
|
||||||
# can safely be reused across multiple agent.run() calls.
|
# and reuse across sequential agent.run() calls.
|
||||||
search_count = 0
|
search_counts: dict[str, int] = {}
|
||||||
current_run_id: str | None = None
|
|
||||||
|
|
||||||
async def search(
|
async def search(
|
||||||
ctx: RunContext[RAGDeps],
|
ctx: RunContext[RAGDeps],
|
||||||
|
|
@ -51,12 +50,9 @@ def create_search_toolset(
|
||||||
Returns:
|
Returns:
|
||||||
Formatted search results with content and metadata.
|
Formatted search results with content and metadata.
|
||||||
"""
|
"""
|
||||||
nonlocal search_count, current_run_id
|
rid = ctx.run_id or ""
|
||||||
if ctx.run_id != current_run_id:
|
search_counts[rid] = search_counts.get(rid, 0) + 1
|
||||||
current_run_id = ctx.run_id
|
if max_searches is not None and search_counts[rid] > max_searches:
|
||||||
search_count = 0
|
|
||||||
search_count += 1
|
|
||||||
if max_searches is not None and search_count > max_searches:
|
|
||||||
return (
|
return (
|
||||||
"Search limit reached. "
|
"Search limit reached. "
|
||||||
"Answer the question using the results you already have."
|
"Answer the question using the results you already have."
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue