add more logging
This commit is contained in:
parent
7f31166829
commit
ff67515b0e
2 changed files with 64 additions and 3 deletions
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue