Files
aisou/ai_search_agent/agent.py
T
TetrisGGBOBot 6ef52a410c feat: add fast/deep time budgets and COT streaming events
Document fast/deep search modes in README, enforce strict 30s/120s timeout behavior with 504 responses, and cap Jina search hits to 10 while exposing richer SSE process events for frontend interaction.

Made-with: Cursor
2026-03-31 12:21:17 +00:00

442 lines
18 KiB
Python

from __future__ import annotations
from collections.abc import AsyncGenerator
from datetime import datetime, timezone
from .cache import SearchCache
from .config import Settings
from .jina import JinaClient
from .knowledge_graph import KnowledgeGraphExtractor
from .llm_client import LLMClient
from .models import (
AnswerPayload,
Citation,
KnowledgeTriple,
PageContent,
SearchRequest,
SearchResponse,
SearchResult,
SearchRoundTrace,
)
from .parsers import get_domain_trust
from .planner import QueryPlanner
from .prompts import SYSTEM_PROMPT
from .reranker import JinaReranker
from .search_client import SearchClient
class AISearchAgent:
FAST_TIMEOUT_SECONDS = 30
DEEP_TIMEOUT_SECONDS = 120
MAX_SEARCH_HITS = 10
def __init__(self, settings: Settings):
self.settings = settings
self.jina = JinaClient(settings)
self.llm = LLMClient(settings)
self.search = SearchClient(settings)
self.cache = SearchCache(settings)
self.reranker = JinaReranker(settings)
self.planner = QueryPlanner(settings)
self.kg = KnowledgeGraphExtractor(settings)
# ------------------------------------------------------------------ #
# Public interface #
# ------------------------------------------------------------------ #
async def run(self, request: SearchRequest) -> SearchResponse:
result: SearchResponse | None = None
async for event in self.run_stream(request):
if event["type"] == "result":
result = SearchResponse.model_validate(event["data"])
if result is None:
raise RuntimeError("Search run completed without result payload")
return result
async def run_stream(
self, request: SearchRequest
) -> AsyncGenerator[dict, None]:
mode = self.resolve_search_mode(request)
timeout_seconds = self.get_time_budget_seconds(request)
max_rounds = 1 if mode == "fast" else self.settings.max_search_rounds
top_k = request.top_k_pages or self.settings.top_k_pages
max_chars = request.max_page_chars or self.settings.max_page_chars
cache_query = f"{mode}:{request.query}"
yield {
"type": "mode_budget",
"data": {
"search_mode": mode,
"deadline_sec": timeout_seconds,
"max_rounds": max_rounds,
"max_hits": self.MAX_SEARCH_HITS,
},
}
yield {
"type": "progress",
"data": {
"stage": "planning",
"message": "开始分析查询并规划搜索。",
"progress": 5,
},
}
# ── 0. Cache hit ────────────────────────────────────────────────
if self.settings.enable_cache:
cached = await self.cache.get_search(cache_query, top_k, max_chars)
if cached:
cached["cache_hit"] = True
cached["search_mode"] = mode
yield {"type": "cache_hit", "data": {"query": request.query}}
yield {"type": "result", "data": cached}
return
# ── 1. History memory ────────────────────────────────────────────
history_context = ""
if self.settings.enable_memory:
history = await self.cache.get_history(5)
recent = [h.get("query", "") for h in history if h.get("query")]
if recent:
history_context = "User's recent searches: " + "; ".join(recent[:5])
# ── 2. Query decomposition ───────────────────────────────────────
sub_queries = await self.planner.decompose(request.query)
yield {"type": "decomposed", "data": {"sub_queries": sub_queries}}
# ── 3. Multilingual expansion ────────────────────────────────────
expanded: list[str] = []
seen_q: set[str] = set()
for sq in sub_queries:
for q in await self.planner.expand_multilingual(sq):
if q not in seen_q:
seen_q.add(q)
expanded.append(q)
if len(expanded) > len(sub_queries):
yield {"type": "expanded", "data": {"queries": expanded}}
yield {
"type": "query_plan",
"data": {
"original_query": request.query,
"sub_queries": sub_queries,
"expanded_queries": expanded,
},
}
# ── 4. Multi-round search loop ───────────────────────────────────
all_results: list[SearchResult] = []
all_pages: list[PageContent] = []
all_usable: list[PageContent] = []
round_traces: list[SearchRoundTrace] = []
latest_digest = ""
seen_urls: set[str] = set()
current_queries = expanded
round_index = 0
for round_index in range(1, max_rounds + 1):
yield {
"type": "round_started",
"data": {"round_index": round_index, "queries": current_queries},
}
yield {
"type": "progress",
"data": {
"stage": "searching",
"message": f"第 {round_index} 轮搜索中。",
"progress": min(20 + round_index * 15, 55),
"round_index": round_index,
},
}
round_results: list[SearchResult] = []
for q in current_queries:
try:
raw_text, results = await self.search.search(q, limit=self.MAX_SEARCH_HITS)
latest_digest = raw_text
for r in results:
if r.url not in seen_urls:
seen_urls.add(r.url)
round_results.append(r)
except Exception as exc:
yield {"type": "search_error", "data": {"query": q, "error": str(exc)}}
yield {
"type": "search_hits",
"data": {
"round_index": round_index,
"queries": current_queries,
"hits": [item.model_dump() for item in round_results[: self.MAX_SEARCH_HITS]],
},
}
pages = await self.jina.read_pages(
round_results[: self.MAX_SEARCH_HITS], max_page_chars=max_chars
)
usable = [p for p in pages if p.fetched and p.usable]
yield {
"type": "page_fetch_status",
"data": {
"round_index": round_index,
"total": len(pages),
"done": len(pages),
"success": len([p for p in pages if p.fetched]),
"failed": len([p for p in pages if not p.fetched]),
},
}
if self.settings.enable_reranker and len(usable) > 1:
try:
usable = await self.reranker.rerank(request.query, usable)
except Exception:
pass
reflection = self._reflect(
round_results, pages, usable, round_index, max_rounds,
len(all_usable) + len(usable),
)
trace = SearchRoundTrace(
round_index=round_index,
query=current_queries[0] if current_queries else request.query,
queries=current_queries,
result_count=len(round_results),
fetched_page_count=len([p for p in pages if p.fetched]),
usable_page_count=len(usable),
reflection=reflection,
)
round_traces.append(trace)
all_results.extend(round_results)
all_pages.extend(pages)
all_usable.extend(usable)
yield {"type": "round_finished", "data": trace.model_dump()}
yield {
"type": "reasoning_note",
"data": {
"round_index": round_index,
"reflection": reflection,
},
}
if len(all_usable) >= top_k or round_index >= max_rounds:
break
next_q = await self.planner.next_query(
request.query, round_index, latest_digest, len(all_usable)
)
if not next_q or next_q in seen_q:
break
seen_q.add(next_q)
current_queries = [next_q]
# ── 5. No usable content ─────────────────────────────────────────
if not all_usable:
response = SearchResponse(
query=request.query,
answer=AnswerPayload(
summary="未读取到可用正文内容,已拒绝调用模型总结。",
caveats=[
"所有候选页面都未能提取到有效正文,可能是 403、451、反爬或页面过短。",
"请更换查询词,或增加可抓取来源。",
],
),
search_results=all_results,
pages=all_pages,
search_rounds=round_traces,
content_ready=False,
llm_called=False,
error="NO_USABLE_CONTENT",
)
yield {"type": "result", "data": response.model_dump()}
return
# ── 6. Prepare model inputs ──────────────────────────────────────
pages_for_model = all_usable[:top_k]
results_for_model = [
r for r in all_results if any(p.url == r.url for p in pages_for_model)
]
allowed_urls = {p.url for p in pages_for_model if p.fetched}
user_prompt = self._build_user_prompt(
request.query, latest_digest, results_for_model, pages_for_model, history_context
)
# ── 7. Token-level LLM streaming ─────────────────────────────────
yield {
"type": "llm_started",
"data": {"page_count": len(pages_for_model), "result_count": len(results_for_model)},
}
yield {
"type": "progress",
"data": {
"stage": "answering",
"message": "正在生成最终回答。",
"progress": 80,
},
}
answer: AnswerPayload | None = None
async for chunk, final_answer in self.llm.answer_stream(SYSTEM_PROMPT, user_prompt, allowed_urls):
if final_answer is not None:
answer = final_answer
elif chunk:
yield {"type": "llm_token", "data": {"text": chunk}}
if answer is None:
answer = AnswerPayload(summary="", caveats=["LLM 未返回有效响应。"])
# ── 8. Confidence-based auto re-search ───────────────────────────
if answer.confidence < self.settings.confidence_threshold:
extra_q = await self.planner.next_query(
request.query, round_index + 1, latest_digest, len(all_usable)
)
if extra_q and extra_q not in seen_q:
yield {
"type": "confidence_retry",
"data": {"confidence": answer.confidence, "extra_query": extra_q},
}
try:
_, extra_results = await self.search.search(
extra_q, limit=self.MAX_SEARCH_HITS
)
new_results = [r for r in extra_results if r.url not in seen_urls]
if new_results:
extra_pages = await self.jina.read_pages(
new_results[:3], max_page_chars=max_chars
)
extra_usable = [p for p in extra_pages if p.fetched and p.usable]
if extra_usable:
pages_for_model = (pages_for_model + extra_usable)[: top_k + 2]
allowed_urls.update(p.url for p in extra_usable)
user_prompt2 = self._build_user_prompt(
request.query, latest_digest,
results_for_model, pages_for_model, history_context,
)
answer = await self.llm.answer(SYSTEM_PROMPT, user_prompt2, allowed_urls)
except Exception:
pass
# ── 9. Backfill citations + domain trust scoring ─────────────────
if not answer.citations:
answer = self._backfill_citations(answer, results_for_model)
answer = self._score_citations(answer)
# ── 10. KG extraction ────────────────────────────────────────────
knowledge_triples: list[KnowledgeTriple] = []
if self.settings.enable_kg:
combined_text = "\n".join(p.content for p in pages_for_model if p.content)
try:
knowledge_triples = await self.kg.extract(combined_text, request.query)
except Exception:
pass
# ── 11. Assemble response ────────────────────────────────────────
response = SearchResponse(
query=request.query,
search_mode=mode,
answer=answer,
search_results=results_for_model,
pages=all_pages,
search_rounds=round_traces,
knowledge_triples=knowledge_triples,
content_ready=True,
llm_called=True,
)
# ── 12. Cache + history ──────────────────────────────────────────
if self.settings.enable_cache:
await self.cache.set_search(cache_query, top_k, max_chars, response.model_dump())
if self.settings.enable_memory:
await self.cache.push_history(
{
"query": request.query,
"summary": answer.summary[:100],
"ts": datetime.now(timezone.utc).isoformat(),
}
)
yield {
"type": "final_meta",
"data": {
"search_mode": mode,
"timed_out": False,
"used_rounds": len(round_traces),
"used_hits": len(results_for_model),
},
}
yield {"type": "result", "data": response.model_dump()}
# ------------------------------------------------------------------ #
# Helpers #
# ------------------------------------------------------------------ #
def _reflect(
self,
results: list[SearchResult],
pages: list[PageContent],
usable: list[PageContent],
round_index: int,
max_rounds: int,
total_usable: int,
) -> str:
if not results:
return "未解析到搜索结果,模型将规划新查询词。"
if not pages:
return "未读取到任何页面内容,下一轮扩大来源。"
if total_usable >= self.settings.top_k_pages:
return "可用证据已达阈值,进入模型总结。"
if round_index >= max_rounds:
return "达到最大搜索轮次,使用当前证据总结。"
blocked = sum(1 for p in pages if p.fetched and not p.usable)
if blocked:
return f"本轮 {blocked} 个页面疑似反爬或正文不足,下一轮将换方向。"
return "证据仍不足,模型将规划下一轮查询。"
def _build_user_prompt(
self,
query: str,
raw_search_text: str,
search_results: list[SearchResult],
pages: list[PageContent],
history_context: str = "",
) -> str:
result_lines: list[str] = []
for r in search_results:
result_lines.append(f"[Result {r.rank}] {r.title or 'Untitled'}")
result_lines.append(f"URL: {r.url}")
if r.snippet:
result_lines.append(f"Snippet: {r.snippet}")
page_lines: list[str] = []
for i, page in enumerate(pages, 1):
page_lines.append(f"[Page {i}] {page.title or page.url}")
page_lines.append(f"URL: {page.url}")
page_lines.append(page.content if page.fetched else f"Fetch error: {page.error}")
ctx = f"\nUser context: {history_context}\n" if history_context else ""
return (
f"User query:\n{query}{ctx}\n\n"
f"Search digest:\n{raw_search_text[:2000]}\n\n"
f"Selected search results:\n{chr(10).join(result_lines)}\n\n"
f"Fetched page excerpts:\n{chr(10).join(page_lines)}"
)
def _backfill_citations(
self, answer: AnswerPayload, results: list[SearchResult]
) -> AnswerPayload:
fallback = [Citation(title=r.title or r.url, url=r.url) for r in results[:3]]
return answer.model_copy(update={"citations": fallback})
def _score_citations(self, answer: AnswerPayload) -> AnswerPayload:
scored = [
c.model_copy(update={"trust_score": get_domain_trust(c.url)})
for c in answer.citations
]
return answer.model_copy(update={"citations": scored})
def resolve_search_mode(self, request: SearchRequest) -> str:
return request.search_mode if request.search_mode in {"fast", "deep"} else "fast"
def get_time_budget_seconds(self, request: SearchRequest) -> int:
return (
self.DEEP_TIMEOUT_SECONDS
if self.resolve_search_mode(request) == "deep"
else self.FAST_TIMEOUT_SECONDS
)