add more logging

This commit is contained in:
bryan davis 2026-04-06 14:10:22 -05:00
parent 7f31166829
commit ff67515b0e
2 changed files with 64 additions and 3 deletions

View file

@ -1,3 +1,5 @@
import logging
import time
from dataclasses import dataclass from dataclasses import dataclass
from pydantic_ai import Agent from pydantic_ai import Agent
@ -15,6 +17,8 @@ from haiku.rag.store.models import SearchResult
from haiku.rag.tools.search import create_search_toolset from haiku.rag.tools.search import create_search_toolset
from haiku.rag.utils import get_model from haiku.rag.utils import get_model
logger = logging.getLogger(__name__)
@dataclass @dataclass
class _QARunDeps: class _QARunDeps:
@ -73,8 +77,23 @@ class QuestionAnswerAgent:
) )
deps = _QARunDeps(client=self._client) deps = _QARunDeps(client=self._client)
result = await agent.run(question, deps=deps)
output = result.output
t0 = time.perf_counter()
result = await agent.run(question, deps=deps)
agent_duration = time.perf_counter() - t0
logger.info(
"qa.agent_run took %.3fs", agent_duration
)
t0 = time.perf_counter()
output = result.output
citations = resolve_citations(output.cited_chunks, accumulated_results) citations = resolve_citations(output.cited_chunks, accumulated_results)
logger.info(
"qa.resolve_citations count=%d took %.3fs",
len(citations),
time.perf_counter() - t0,
)
logger.info(
"qa.answer completed total=%.3fs", agent_duration
)
return output.answer, citations return output.answer, citations

View file

@ -1,3 +1,5 @@
import logging
import time
from collections.abc import Callable from collections.abc import Callable
from pydantic_ai import FunctionToolset, RunContext from pydantic_ai import FunctionToolset, RunContext
@ -6,6 +8,8 @@ from haiku.rag.config.models import AppConfig
from haiku.rag.store.models import SearchResult from haiku.rag.store.models import SearchResult
from haiku.rag.tools.context import RAGDeps from haiku.rag.tools.context import RAGDeps
logger = logging.getLogger(__name__)
def create_search_toolset( def create_search_toolset(
config: AppConfig, config: AppConfig,
@ -35,6 +39,7 @@ def create_search_toolset(
# Per-run search counter keyed by run_id. Safe for concurrent runs # Per-run search counter keyed by run_id. Safe for concurrent runs
# and reuse across sequential agent.run() calls. # and reuse across sequential agent.run() calls.
search_counts: dict[str, int] = {} search_counts: dict[str, int] = {}
_last_tool_return: list[float] = []
async def search( async def search(
ctx: RunContext[RAGDeps], ctx: RunContext[RAGDeps],
@ -50,9 +55,18 @@ def create_search_toolset(
Returns: Returns:
Formatted search results with content and metadata. Formatted search results with content and metadata.
""" """
tool_start = time.perf_counter()
if _last_tool_return:
llm_think_time = tool_start - _last_tool_return[0]
logger.info(
"tool.llm_thinking took %.3fs", llm_think_time
)
rid = ctx.run_id or "" rid = ctx.run_id or ""
search_counts[rid] = search_counts.get(rid, 0) + 1 search_counts[rid] = search_counts.get(rid, 0) + 1
if max_searches is not None and search_counts[rid] > max_searches: if max_searches is not None and search_counts[rid] > max_searches:
_last_tool_return[:] = [time.perf_counter()]
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."
@ -62,12 +76,25 @@ def create_search_toolset(
effective_filter = base_filter effective_filter = base_filter
effective_limit = limit or config.search.limit effective_limit = limit or config.search.limit
t0 = time.perf_counter()
results = await client.search( results = await client.search(
query, limit=effective_limit, filter=effective_filter query, limit=effective_limit, filter=effective_filter
) )
logger.info(
"tool.search query=%r took %.3fs",
query[:80],
time.perf_counter() - t0,
)
if expand_context: if expand_context:
t0 = time.perf_counter()
results = await client.expand_context(results) results = await client.expand_context(results)
logger.info(
"tool.expand_context results=%d took %.3fs",
len(results),
time.perf_counter() - t0,
)
results_list = list(results) results_list = list(results)
@ -75,14 +102,29 @@ def create_search_toolset(
on_results(results_list) on_results(results_list)
if not results_list: if not results_list:
_last_tool_return[:] = [time.perf_counter()]
return "No results found." return "No results found."
t0 = time.perf_counter()
total = len(results_list) total = len(results_list)
formatted = [ formatted = [
r.format_for_agent(rank=i + 1, total=total) r.format_for_agent(rank=i + 1, total=total)
for i, r in enumerate(results_list) for i, r in enumerate(results_list)
] ]
return "\n\n".join(formatted) output = "\n\n".join(formatted)
logger.info(
"tool.format results=%d chars=%d took %.3fs",
total,
len(output),
time.perf_counter() - t0,
)
logger.info(
"tool.search_total took %.3fs",
time.perf_counter() - tool_start,
)
_last_tool_return[:] = [time.perf_counter()]
return output
toolset: FunctionToolset[RAGDeps] = FunctionToolset() toolset: FunctionToolset[RAGDeps] = FunctionToolset()
toolset.add_function(search, name=tool_name, retries=3) toolset.add_function(search, name=tool_name, retries=3)