Refactor knowledge search handler to use multiprocessing for isolated context retrieval and improved stability.
This commit is contained in:
@@ -1,6 +1,8 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import multiprocessing
|
||||||
|
import time
|
||||||
from typing import Any, Dict
|
from typing import Any, Dict
|
||||||
|
|
||||||
from ...local_rag import LibraryContextRequest, library_context
|
from ...local_rag import LibraryContextRequest, library_context
|
||||||
@@ -22,13 +24,51 @@ def _knowledge_source(source: Dict[str, Any], library_slug: str) -> Dict[str, An
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
async def knowledge_search_handler(arguments: Dict[str, Any], _context: ToolExecutionContext) -> Dict[str, Any]:
|
def _knowledge_search_worker(connection, arguments: Dict[str, Any]) -> None:
|
||||||
request = LibraryContextRequest(
|
try:
|
||||||
prompt=arguments["prompt"],
|
request = LibraryContextRequest(
|
||||||
top_k=arguments.get("top_k", 5),
|
prompt=arguments["prompt"],
|
||||||
embed_model=arguments.get("embedding_model"),
|
top_k=arguments.get("top_k", 5),
|
||||||
)
|
embed_model=arguments.get("embedding_model"),
|
||||||
payload = await asyncio.to_thread(library_context, arguments["library_slug"], request)
|
)
|
||||||
|
connection.send(("ok", library_context(arguments["library_slug"], request)))
|
||||||
|
except BaseException as exc:
|
||||||
|
detail = getattr(exc, "detail", None) or str(exc)
|
||||||
|
connection.send(("error", f"{type(exc).__name__}: {detail}"))
|
||||||
|
finally:
|
||||||
|
connection.close()
|
||||||
|
|
||||||
|
|
||||||
|
async def _isolated_library_context(arguments: Dict[str, Any], context: ToolExecutionContext) -> Dict[str, Any]:
|
||||||
|
process_context = multiprocessing.get_context("spawn")
|
||||||
|
parent_connection, child_connection = process_context.Pipe(duplex=False)
|
||||||
|
process = process_context.Process(target=_knowledge_search_worker, args=(child_connection, arguments), daemon=True)
|
||||||
|
process.start()
|
||||||
|
child_connection.close()
|
||||||
|
deadline = time.monotonic() + 170
|
||||||
|
try:
|
||||||
|
while True:
|
||||||
|
if context.cancellation_event.is_set():
|
||||||
|
raise asyncio.CancelledError()
|
||||||
|
if parent_connection.poll():
|
||||||
|
status, value = parent_connection.recv()
|
||||||
|
if status == "ok":
|
||||||
|
return value
|
||||||
|
raise RuntimeError(value)
|
||||||
|
if not process.is_alive():
|
||||||
|
raise RuntimeError(f"Local retrieval process exited unexpectedly (exit code {process.exitcode}).")
|
||||||
|
if time.monotonic() >= deadline:
|
||||||
|
raise TimeoutError("Local retrieval exceeded its process timeout.")
|
||||||
|
await asyncio.sleep(0.05)
|
||||||
|
finally:
|
||||||
|
parent_connection.close()
|
||||||
|
if process.is_alive():
|
||||||
|
process.terminate()
|
||||||
|
process.join(timeout=2)
|
||||||
|
|
||||||
|
|
||||||
|
async def knowledge_search_handler(arguments: Dict[str, Any], context: ToolExecutionContext) -> Dict[str, Any]:
|
||||||
|
payload = await _isolated_library_context(arguments, context)
|
||||||
raw_hits = (payload.get("result") or {}).get("sources") or []
|
raw_hits = (payload.get("result") or {}).get("sources") or []
|
||||||
budget = int(arguments.get("context_character_budget") or 12_000)
|
budget = int(arguments.get("context_character_budget") or 12_000)
|
||||||
context_block = str(payload.get("context_block") or "")[:budget]
|
context_block = str(payload.get("context_block") or "")[:budget]
|
||||||
|
|||||||
Reference in New Issue
Block a user