|
# retoor <retoor@molodetz.nl>
|
|
|
|
import unittest
|
|
from collections.abc import Callable
|
|
|
|
from typosaurus_sandbox.research.config import ResearchConfig
|
|
from typosaurus_sandbox.research.engine import ResearchEngine
|
|
from typosaurus_sandbox.research.envelopes import ChatResponse, DescribeResponse, SearchResponse, SearchResult
|
|
|
|
|
|
class FakeResearchClient:
|
|
def __init__(
|
|
self,
|
|
*,
|
|
web_results: list[SearchResult] | None = None,
|
|
web_result_factory: Callable[[str], list[SearchResult]] | None = None,
|
|
chat_text: str = "",
|
|
describe_text: str = "",
|
|
) -> None:
|
|
self.config = ResearchConfig(max_concurrency=4, default_count=5)
|
|
self._web_results = web_results if web_results is not None else []
|
|
self._web_result_factory = web_result_factory
|
|
self._chat_text = chat_text
|
|
self._describe_text = describe_text
|
|
self.calls: list[tuple[str, str, str | None]] = []
|
|
|
|
def search_cached(
|
|
self,
|
|
query: str,
|
|
*,
|
|
source: str | None = None,
|
|
count: int | None = None,
|
|
content: bool = False,
|
|
type: str | None = None,
|
|
deep: bool = False,
|
|
ai: bool = False,
|
|
cache: bool = True,
|
|
) -> None:
|
|
return None
|
|
|
|
def describe_cached(self, url: str) -> None:
|
|
return None
|
|
|
|
async def search(
|
|
self,
|
|
query: str,
|
|
*,
|
|
source: str | None = None,
|
|
count: int | None = None,
|
|
content: bool = False,
|
|
type: str | None = None,
|
|
deep: bool = False,
|
|
ai: bool = False,
|
|
cache: bool = True,
|
|
) -> SearchResponse:
|
|
self.calls.append(("search", query, type))
|
|
if type == "images":
|
|
return SearchResponse(query=query, source="wikimedia", count=0, success=True, results=[])
|
|
results = self._web_result_factory(query) if self._web_result_factory is not None else list(self._web_results)
|
|
return SearchResponse(query=query, source="duckduckgo", count=len(results), success=True, results=results)
|
|
|
|
async def chat(self, prompt: str, *, system: str | None = None, json_mode: bool = False, cache: bool = True) -> ChatResponse:
|
|
self.calls.append(("chat", prompt, None))
|
|
return ChatResponse(response=self._chat_text, prompt=prompt)
|
|
|
|
async def describe(self, url: str) -> DescribeResponse:
|
|
self.calls.append(("describe", url, None))
|
|
return DescribeResponse(description=self._describe_text, url=url)
|
|
|
|
|
|
class TestEngineClosureDetection(unittest.IsolatedAsyncioTestCase):
|
|
|
|
async def test_run_closes_after_single_round_when_nothing_new(self) -> None:
|
|
client = FakeResearchClient(web_results=[], chat_text="", describe_text="")
|
|
engine = ResearchEngine(client=client)
|
|
report = await engine.run(" deep research ")
|
|
self.assertEqual(report.subject, "deep research")
|
|
self.assertTrue(report.closed)
|
|
self.assertEqual(report.total_rounds, 1)
|
|
self.assertEqual(len(report.rounds), 1)
|
|
first = report.rounds[0]
|
|
self.assertEqual(first.number, 1)
|
|
self.assertEqual(first.items_processed, 3)
|
|
self.assertEqual(first.requests_succeeded, 3)
|
|
self.assertEqual(first.requests_failed, 0)
|
|
self.assertEqual(first.new_urls, 0)
|
|
self.assertEqual(first.new_queries, 0)
|
|
self.assertTrue(first.closed)
|
|
self.assertEqual(report.queries_issued, 1)
|
|
self.assertEqual(report.queries_enqueued, 1)
|
|
self.assertEqual(report.urls_collected, 0)
|
|
self.assertEqual(report.contents_seen, 0)
|
|
self.assertEqual(report.content_types, {"web": 1, "images": 1, "chat": 1})
|
|
self.assertEqual(len(client.calls), 3)
|
|
self.assertEqual({kind for kind, _, _ in client.calls}, {"search", "chat"})
|
|
|
|
async def test_run_discovery_rounds_then_closes(self) -> None:
|
|
client = FakeResearchClient(
|
|
web_results=[
|
|
SearchResult(
|
|
title="topic alpha",
|
|
url="https://example.com/alpha",
|
|
description="alpha details",
|
|
content="alpha body",
|
|
)
|
|
],
|
|
chat_text="",
|
|
describe_text="",
|
|
)
|
|
engine = ResearchEngine(client=client)
|
|
report = await engine.run("deep research")
|
|
self.assertTrue(report.closed)
|
|
self.assertEqual(report.total_rounds, 2)
|
|
first = report.rounds[0]
|
|
self.assertEqual(first.new_urls, 1)
|
|
self.assertEqual(first.new_queries, 2)
|
|
self.assertEqual(first.new_contents, 1)
|
|
self.assertFalse(first.closed)
|
|
second = report.rounds[1]
|
|
self.assertEqual(second.new_urls, 0)
|
|
self.assertEqual(second.new_queries, 0)
|
|
self.assertTrue(second.closed)
|
|
self.assertEqual(report.queries_generated, 7)
|
|
self.assertEqual(report.queries_enqueued, 3)
|
|
self.assertEqual(report.queries_issued, 3)
|
|
self.assertEqual(report.queries_duplicates_skipped, 4)
|
|
self.assertEqual(report.urls_collected, 1)
|
|
self.assertEqual(report.urls_duplicates_skipped, 2)
|
|
self.assertEqual(report.contents_seen, 1)
|
|
self.assertEqual(report.content_duplicates_skipped, 2)
|
|
self.assertEqual(report.requests_succeeded, 10)
|
|
self.assertEqual(report.requests_failed, 0)
|
|
self.assertEqual(report.cache_hits, 0)
|
|
self.assertEqual(report.cache_misses, 10)
|
|
self.assertEqual(report.content_types, {"web": 3, "images": 3, "chat": 3, "describe": 1})
|
|
self.assertEqual(len(client.calls), 10)
|
|
self.assertEqual(sum(1 for kind, _, type_value in client.calls if kind == "search" and type_value is None), 3)
|
|
self.assertEqual(sum(1 for kind, _, type_value in client.calls if kind == "search" and type_value == "images"), 3)
|
|
self.assertEqual(sum(1 for kind, _, _ in client.calls if kind == "chat"), 3)
|
|
self.assertEqual(sum(1 for kind, value, _ in client.calls if kind == "describe"), 1)
|
|
self.assertIn(("describe", "https://example.com/alpha", None), client.calls)
|
|
|
|
async def test_new_content_alone_does_not_prevent_closure(self) -> None:
|
|
def factory(query: str) -> list[SearchResult]:
|
|
return [
|
|
SearchResult(
|
|
title="dup title",
|
|
url="https://example.com/dup",
|
|
description="dup details",
|
|
content=f"body for {query}",
|
|
)
|
|
]
|
|
|
|
client = FakeResearchClient(web_result_factory=factory, chat_text="", describe_text="")
|
|
engine = ResearchEngine(client=client)
|
|
report = await engine.run("subject")
|
|
self.assertTrue(report.closed)
|
|
self.assertEqual(report.total_rounds, 2)
|
|
first = report.rounds[0]
|
|
self.assertEqual(first.new_urls, 1)
|
|
self.assertEqual(first.new_queries, 2)
|
|
self.assertFalse(first.closed)
|
|
second = report.rounds[1]
|
|
self.assertEqual(second.new_urls, 0)
|
|
self.assertEqual(second.new_queries, 0)
|
|
self.assertEqual(second.new_contents, 2)
|
|
self.assertTrue(second.closed)
|
|
self.assertEqual(report.contents_seen, 3)
|
|
|
|
async def test_round_summary_dict_is_serialisable(self) -> None:
|
|
client = FakeResearchClient(web_results=[], chat_text="", describe_text="")
|
|
engine = ResearchEngine(client=client)
|
|
report = await engine.run("serialisable subject")
|
|
summary_dict = report.rounds[0].to_dict()
|
|
self.assertEqual(summary_dict["number"], 1)
|
|
self.assertTrue(summary_dict["closed"])
|
|
report_dict = report.to_dict()
|
|
self.assertEqual(report_dict["subject"], "serialisable subject")
|
|
self.assertEqual(report_dict["total_rounds"], 1)
|
|
self.assertTrue(report_dict["closed"])
|
|
|
|
|
|
class TestEngineInputValidation(unittest.IsolatedAsyncioTestCase):
|
|
|
|
async def test_empty_subject_raises_without_requests(self) -> None:
|
|
client = FakeResearchClient(web_results=[], chat_text="", describe_text="")
|
|
engine = ResearchEngine(client=client)
|
|
with self.assertRaises(ValueError) as ctx:
|
|
await engine.run(" \n\t ")
|
|
self.assertEqual(str(ctx.exception), "research subject must not be empty")
|
|
self.assertEqual(client.calls, [])
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|
|
|
|
|
|
|