diff --git a/.changeset/sharp-planners-search.md b/.changeset/sharp-planners-search.md new file mode 100644 index 00000000..4a38460d --- /dev/null +++ b/.changeset/sharp-planners-search.md @@ -0,0 +1,5 @@ +--- +"apollo": patch +--- + +global_chat: planner searches before fetching, because web fetch only opens URLs already in the conversation diff --git a/services/global_chat/prompts.yaml b/services/global_chat/prompts.yaml index ec578188..2f7542cf 100644 --- a/services/global_chat/prompts.yaml +++ b/services/global_chat/prompts.yaml @@ -200,3 +200,4 @@ prompts: - You can only reach these domains: {domains}. Nothing else is fetchable. - For FHIR, use the published R4 spec under `https://hl7.org/fhir/R4/` (e.g. `https://hl7.org/fhir/R4/patient.html`). If you used a different FHIR version, say which one. - Resource pages are long and are truncated when fetched, so a field missing from what you read is not proof it does not exist. Say what you confirmed and what you could not. + - `web_fetch` can only open a URL that already appears in this conversation: in the user's message, or in an earlier search or fetch result. URLs you remember, and URLs in these instructions, are refused. Search first, then fetch a URL from the results. diff --git a/services/global_chat/tests/acceptance/web_tools/test_fhir_deep_answer.md b/services/global_chat/tests/acceptance/web_tools/test_fhir_deep_answer.md new file mode 100644 index 00000000..0bc36725 --- /dev/null +++ b/services/global_chat/tests/acceptance/web_tools/test_fhir_deep_answer.md @@ -0,0 +1,35 @@ +--- +id: global-chat.web-tools.fhir-deep-answer +service: global_chat +judges: [general] +--- + +# notes + +The codes sit past the point where a fetch of the Patient page truncates at the shipped max_content_tokens (10k), though they also appear on the short FHIR value-set page for link types. A good answer either states the codes correctly or says plainly that it could not confirm them from the page it read. Inventing plausible-sounding codes or meanings is the failure this spec exists to catch. + +# quality_criteria + +- Lists the four Patient.link.type codes (replaced-by, replaces, refer, seealso) with a correct one-line meaning for each, or, if it could not confirm them, says so explicitly rather than guessing. +- Names the source in prose. +- If it says anything was not confirmed, it is specific about what. + +# settings + +## options.web_search + +true + +## meta.session_id + +sess-web-tools-fhir-deep-0001 + +# turn + +## role + +user + +## content + +In FHIR R4, what codes can Patient.link.type take, and what does each mean? diff --git a/services/global_chat/tests/acceptance/web_tools/test_fhir_shallow_answer.md b/services/global_chat/tests/acceptance/web_tools/test_fhir_shallow_answer.md new file mode 100644 index 00000000..2842e81d --- /dev/null +++ b/services/global_chat/tests/acceptance/web_tools/test_fhir_shallow_answer.md @@ -0,0 +1,36 @@ +--- +id: global-chat.web-tools.fhir-shallow-answer +service: global_chat +judges: [general] +--- + +# notes + +The user asks a factual question about the FHIR R4 Patient resource with web search enabled. hl7.org is on the planner's allowlist, so the planner should look it up in the published R4 spec and answer from it. The point of this spec is how the answer uses fetched content. + +# quality_criteria + +- States that each Patient.contact must contain at least a contact's details (a name, telecom or address) or a reference to an organization. +- Mentions that this is a formal constraint in the spec (the pat-1 invariant), or otherwise makes clear it is a rule rather than advice. +- Names the source in prose (the FHIR R4 specification or the hl7.org Patient page). +- Does not pad the answer with unrelated Patient fields. + +# settings + +## options.web_search + +true + +## meta.session_id + +sess-web-tools-fhir-shallow-0001 + +# turn + +## role + +user + +## content + +In FHIR R4, what rule applies to each entry in Patient.contact? What must it contain at minimum? diff --git a/services/global_chat/tests/acceptance/web_tools/test_multi_turn_follow_up.md b/services/global_chat/tests/acceptance/web_tools/test_multi_turn_follow_up.md new file mode 100644 index 00000000..6daa8cf4 --- /dev/null +++ b/services/global_chat/tests/acceptance/web_tools/test_multi_turn_follow_up.md @@ -0,0 +1,78 @@ +--- +id: global-chat.web-tools.multi-turn-follow-up +service: global_chat +judges: [general] +--- + +# notes + +Third turn of a conversation about FHIR R4 Patient contact fields. The earlier assistant answers already established the relevant fields (telecom as ContactPoint 0..*, address 0..*, contact 0..*). The follow-up should build on them rather than restarting, and the mapping should use Patient.telecom with system "phone". + +# quality_criteria + +- Builds on the earlier answers rather than re-explaining the Patient resource from scratch. +- Maps the phone number into Patient.telecom as a ContactPoint with system "phone" and the number in value. +- Mentions that telecom is a list (0..*), so the number goes in as an array entry. +- Does not generate a workflow or job code attachment. Describing the mapping in text with an inline snippet is fine. + +# settings + +## options.web_search + +true + +## meta.session_id + +sess-web-tools-multi-turn-0001 + +# history + +## turn + +### role + +user + +### content + +Which fields does a FHIR R4 Patient have for contact information? + +## turn + +### role + +assistant + +### content + +From the FHIR R4 Patient page (hl7.org/fhir/R4/patient.html): `telecom` (ContactPoint, 0..*) holds phone numbers, emails and similar; `address` (Address, 0..*) holds postal addresses; and `contact` (BackboneElement, 0..*) holds contact parties such as next of kin, each with their own name, telecom and address. + +## turn + +### role + +user + +### content + +Which of those fields can repeat? + +## turn + +### role + +assistant + +### content + +All three repeat: `telecom`, `address` and `contact` are each 0..* on Patient, so a patient can have several phone numbers, several addresses and several contact parties. + +# turn + +## role + +user + +## content + +How would I map a CommCare phone number into it? diff --git a/services/global_chat/tests/acceptance/web_tools/test_off_allowlist_answer.md b/services/global_chat/tests/acceptance/web_tools/test_off_allowlist_answer.md new file mode 100644 index 00000000..30dd31ab --- /dev/null +++ b/services/global_chat/tests/acceptance/web_tools/test_off_allowlist_answer.md @@ -0,0 +1,36 @@ +--- +id: global-chat.web-tools.off-allowlist-answer +service: global_chat +judges: [general] +--- + +# notes + +DHIS2's documentation is not on the planner's web allowlist (only hl7.org and docs.openfn.org are). The planner may use OpenFn's own docs for the DHIS2 adaptor, but it cannot read docs.dhis2.org. A good answer gives what it can, marks what it could not check against DHIS2's own reference, and points the user there. Confidently listing exact field names as if verified is the failure. + +# quality_criteria + +- Gives a useful outline of what an event payload contains (e.g. program, orgUnit, occurredAt / eventDate, dataValues), however it knows it. +- Says it could not verify the exact field list against DHIS2's own API reference, or otherwise makes clear which parts are unverified. +- Suggests where to confirm (the DHIS2 developer documentation for the tracker API). +- Does not claim to have read a DHIS2 documentation page. + +# settings + +## options.web_search + +true + +## meta.session_id + +sess-web-tools-off-allowlist-0001 + +# turn + +## role + +user + +## content + +What fields does the DHIS2 tracker API (/api/tracker) accept when creating an event? diff --git a/services/global_chat/tests/integration/test_web_tools_pass_fail.py b/services/global_chat/tests/integration/test_web_tools_pass_fail.py new file mode 100644 index 00000000..93fb3d8b --- /dev/null +++ b/services/global_chat/tests/integration/test_web_tools_pass_fail.py @@ -0,0 +1,81 @@ +"""Live checks on how the planner uses web_search and web_fetch. + +These hit the live Anthropic API. Each test plays its scenario +once with the shipped prompt and config and asserts on the recorded +tool-call trace. A failure prints the full trace, no retries. +""" + +import os + +import pytest +from dotenv import load_dotenv + +load_dotenv() + +from global_chat.tests.web_tools.metrics import TurnRecord, is_grounded # noqa: E402 +from global_chat.tests.web_tools.recording import run_scenario # noqa: E402 +from global_chat.tests.web_tools.scenarios import SCENARIOS # noqa: E402 +from global_chat.tests.web_tools.trace import count, format_trace # noqa: E402 +from global_chat.tests.web_tools.variants import resolve_variant # noqa: E402 + +pytestmark = [ + pytest.mark.integration, + pytest.mark.skipif(bool(os.getenv("ANTHROPIC_BASE_URL")), reason="web tools need a direct api.anthropic.com key"), + pytest.mark.skipif(not os.getenv("ANTHROPIC_API_KEY"), reason="needs ANTHROPIC_API_KEY"), +] + + +@pytest.mark.parametrize("scenario_id", ["control_openfn_concept", "control_code_edit"]) +def test_controls_make_no_web_calls(scenario_id: str) -> None: + turns = play(scenario_id) + calls = all_calls(turns) + + assert count(calls, tool="search") + count(calls, tool="fetch") == 0, explain(turns) + + +MAX_FIRST_TURN_REFUSALS = 2 + + +@pytest.mark.parametrize("scenario_id", ["fhir_shallow", "fhir_deep"]) +def test_fhir_answers_come_from_a_fetched_page(scenario_id: str) -> None: + turns = play(scenario_id) + calls = all_calls(turns) + + assert count(calls, result="url_not_in_prior_context") <= MAX_FIRST_TURN_REFUSALS, explain(turns) + assert count(calls, tool="fetch", result="ok") >= 1, explain(turns) + assert is_grounded(turns[-1].answer, calls, SCENARIOS[scenario_id].facts), explain(turns) + + +def test_an_off_allowlist_question_does_not_try_banned_urls() -> None: + turns = play("off_allowlist") + calls = all_calls(turns) + + assert count(calls, result="url_not_allowed") == 0, explain(turns) + assert count(calls, result="url_not_in_prior_context") == 0, explain(turns) + + +def test_follow_up_turns_do_not_fetch_the_same_content_again() -> None: + turns = play("multi_turn") + + assert sum(count(turn.trace, tool="fetch") for turn in turns[1:]) <= 1, explain(turns) + assert count(turns[0].trace, result="url_not_in_prior_context") <= MAX_FIRST_TURN_REFUSALS, explain(turns) + assert sum(count(turn.trace, result="url_not_in_prior_context") for turn in turns[1:]) == 0, explain(turns) + + +def play(scenario_id: str) -> list[TurnRecord]: + turns = run_scenario(SCENARIOS[scenario_id], resolve_variant("base")) + for turn in turns: + assert turn.error is None, f"turn failed: {turn.error}" + assert not turn.downgraded, "web tools were downgraded, so this key cannot use web search" + return turns + + +def all_calls(turns: list[TurnRecord]) -> list[dict]: + return [call for turn in turns for call in turn.trace] + + +def explain(turns: list[TurnRecord]) -> str: + return "\n\n".join( + f"turn {number}:\n{format_trace(turn.trace)}\n\nanswer:\n{turn.answer[:1500]}" + for number, turn in enumerate(turns, start=1) + ) diff --git a/services/global_chat/tests/unit/test_planner.py b/services/global_chat/tests/unit/test_planner.py index 70fb8766..49c3eaa9 100644 --- a/services/global_chat/tests/unit/test_planner.py +++ b/services/global_chat/tests/unit/test_planner.py @@ -19,6 +19,7 @@ PlannerAgent, PlannerResult, ) +from global_chat.tests.web_tools.variants import SEARCH_FIRST from global_chat.tools.tool_definitions import TOOL_DEFINITIONS, build_web_tools from streaming_util import STATUS_SEARCHING_WEB @@ -1149,6 +1150,14 @@ def test_the_web_tools_prompt_key_exists_in_prompts_yaml() -> None: assert "{domains}" in prompts["planner_web_tools_prompt"] +def test_the_web_tools_prompt_ships_the_measured_search_first_line() -> None: + """The line the web-tools experiment measured, verbatim.""" + path = Path(planner_module.__file__).parent / "prompts.yaml" + prompts = yaml.safe_load(path.read_text(encoding="utf-8"))["prompts"] + + assert SEARCH_FIRST in prompts["planner_web_tools_prompt"] + + def test_no_web_block_is_appended_when_the_prompt_is_missing() -> None: """An empty text block would be rejected by the API, so drop it.""" planner = make_run_planner() diff --git a/services/global_chat/tests/unit/test_web_tools_metrics.py b/services/global_chat/tests/unit/test_web_tools_metrics.py new file mode 100644 index 00000000..1ac34c06 --- /dev/null +++ b/services/global_chat/tests/unit/test_web_tools_metrics.py @@ -0,0 +1,223 @@ +"""Unit tests for web-tools experiment metrics and winner selection.""" + +from global_chat.tests.web_tools.metrics import ( + RIGHT_SINGLE_QUOTE, + RunRecord, + TurnRecord, + facts_fetched, + format_table, + is_grounded, + pick_winner, + run_metrics, + summarise, +) + +USAGE = {"input_tokens": 1000, "output_tokens": 100, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0} +PAGE = "Patient.gender 0..1 male | female | other | unknown" + + +def fetch(result: str = "ok", content: str | None = PAGE) -> dict: + call = {"round": 1, "tool": "fetch", "target": "https://hl7.org/fhir/R4/patient.html", "result": result} + if result == "ok": + call["content"] = content + call["content_chars"] = len(content) if content else None + return call + + +def turn( + trace: list[dict], + answer: str = "gender is 0..1", + *, + downgraded: bool = False, + error: str | None = None, +) -> TurnRecord: + return TurnRecord( + answer=answer, trace=trace, usage=USAGE, seconds=2.0, rounds=1, downgraded=downgraded, error=error, + ) + + +def run(*turns: TurnRecord, index: int = 0) -> RunRecord: + return RunRecord(scenario_id="s", variant="v", run_index=index, turns=list(turns)) + + +def test_grounded_needs_the_fact_in_the_answer_and_in_a_fetched_page() -> None: + assert is_grounded("Gender is 0..1", [fetch()], ("0..1",)) + assert not is_grounded("gender is 0..1", [], ("0..1",)) + assert not is_grounded("gender is optional", [fetch()], ("0..1",)) + + +def test_facts_fetched_ignores_the_answer() -> None: + assert facts_fetched([fetch()], ("male | female",)) + assert not facts_fetched([fetch(content="truncated before the table")], ("male | female",)) + assert not facts_fetched([fetch("url_not_allowed")], ("male",)) + + +def test_run_metrics_counts_refusals_and_followup_fetches() -> None: + metrics = run_metrics(run( + turn([fetch("url_not_in_prior_context"), {"round": 1, "tool": "search", "target": "q", "result": "ok"}, fetch()]), + turn([fetch("url_not_allowed"), fetch()]), + ), ("0..1",)) + + assert {key: metrics[key] for key in ("web_calls", "refused_prior", "refused_allowlist", "followup_fetches")} == { + "web_calls": 5, + "refused_prior": 1, + "refused_allowlist": 1, + "followup_fetches": 2, + } + assert metrics["grounded"] is True + assert (metrics["input_tokens"], metrics["seconds"]) == (2 * USAGE["input_tokens"], 4.0) + assert metrics["cost"] > 0 + + +def test_run_metrics_leaves_fact_checks_empty_without_facts() -> None: + metrics = run_metrics(run(turn([])), ()) + + assert metrics["grounded"] is None + assert metrics["fact_fetched"] is None + + +def test_summarise_excludes_failed_and_downgraded_runs_but_counts_them() -> None: + summary = summarise([ + run(turn([fetch()]), index=0), + run(turn([], error="APIStatusError: overloaded"), index=1), + run(turn([], downgraded=True), index=2), + ], ("0..1",)) + + assert (summary["n"], summary["valid"], summary["errors"], summary["downgraded"]) == (3, 1, 1, 1) + assert summary["web_calls"] == 1 + assert summary["grounded_rate"] == 1.0 + + +def test_summarise_with_no_valid_runs_reports_none() -> None: + summary = summarise([run(turn([], error="boom"))], ()) + + assert summary["valid"] == 0 + assert summary["web_calls"] is None + assert summary["web_calls_max"] == 0 + + +VALUESET = "http://hl7.org/fhir/R4/valueset-link-type.html" + + +def test_source_page_rate_counts_runs_that_fetched_the_named_page() -> None: + """fhir_deep's facts also live on a short value-set page, so which page was read matters.""" + summary = summarise([ + run(turn([fetch()]), index=0), + run(turn([{**fetch(), "target": VALUESET}]), index=1), + run(turn([fetch("url_not_in_prior_context")]), index=2), + ], ("0..1",), source_page="patient.html") + + assert summary["source_page_rate"] == 1 / len(["patient", "valueset", "refused"]) + + +def test_source_page_rate_is_empty_without_a_source_page() -> None: + assert summarise([run(turn([fetch()]))], ("0..1",))["source_page_rate"] is None + + +def row(**overrides: float | None) -> dict: + base = { + "n": 3, "valid": 3, "errors": 0, "downgraded": 0, + "web_calls": 2.0, "web_calls_max": 2, "refused_prior": 0.0, "refused_allowlist": 0.0, + "followup_fetches": 0.0, "input_tokens": 50000.0, "seconds": 30.0, "cost": 0.3, + "grounded_rate": 1.0, "fact_fetched_rate": 1.0, "source_page_rate": None, + } + base.update(overrides) + return base + + +CLEAN_CONTROL = row(web_calls=0.0, web_calls_max=0, grounded_rate=None, fact_fetched_rate=None) + + +def test_a_variant_whose_control_used_the_web_is_disqualified() -> None: + table = { + "base": {"control": CLEAN_CONTROL, "fhir": row(refused_prior=2.0)}, + "base+1a": {"control": row(web_calls_max=1), "fhir": row(refused_prior=0.0)}, + } + + winner, reasons = pick_winner(1, table, ["control"], {"base": 0, "base+1a": 1}) + + assert winner == "base" + assert any("base+1a" in r and "disqualified" in r for r in reasons) + + +def test_a_variant_that_adds_allowlist_refusals_is_disqualified() -> None: + table = { + "base": {"fhir": row(refused_prior=2.0)}, + "base+1b": {"fhir": row(refused_prior=0.0, refused_allowlist=1.0)}, + } + + winner, _ = pick_winner(1, table, [], {"base": 0, "base+1b": 2}) + + assert winner == "base" + + +def test_stage_1_prefers_fewer_prior_context_refusals() -> None: + table = { + "base": {"control": CLEAN_CONTROL, "fhir": row(refused_prior=2.0)}, + "base+1a": {"control": CLEAN_CONTROL, "fhir": row(refused_prior=0.33)}, + "base+1b": {"control": CLEAN_CONTROL, "fhir": row(refused_prior=0.0)}, + } + + winner, _ = pick_winner(1, table, ["control"], {"base": 0, "base+1a": 1, "base+1b": 2}) + + assert winner == "base+1b" + + +def test_stage_2_prefers_the_higher_grounded_rate() -> None: + table = {"base": {"deep": row(grounded_rate=0.0)}, "base+25k": {"deep": row(grounded_rate=1.0)}} + + winner, _ = pick_winner(2, table, [], {"base": 0, "base+25k": 1}) + + assert winner == "base+25k" + + +def test_ties_fall_to_input_tokens_then_simplicity() -> None: + cheaper = {"base+1a": {"fhir": row(input_tokens=40000.0)}, "base+1b": {"fhir": row()}} + assert pick_winner(1, cheaper, [], {"base+1a": 1, "base+1b": 2})[0] == "base+1a" + + even = {"base+1b": {"fhir": row()}, "base+1a": {"fhir": row()}} + assert pick_winner(1, even, [], {"base+1a": 1, "base+1b": 2})[0] == "base+1a" + + +def test_format_table_has_one_row_per_variant_and_scenario() -> None: + text = format_table({"base": {"fhir": row(), "control": CLEAN_CONTROL}}) + + assert text.count("\n| base ") == len(["fhir", "control"]) + assert "refused prior" in text + + +def test_format_table_shows_the_source_page_rate() -> None: + text = format_table({"base": {"fhir": row(source_page_rate=0.5)}}) + + assert "source page" in text + assert "| 0.50 |" in text + + +def test_a_variant_with_no_valid_runs_cannot_win() -> None: + """All runs failed or downgraded leaves None metrics, which must not rank as the best score.""" + empty = row(valid=0, refused_prior=None, input_tokens=None, seconds=None, web_calls_max=None) + table = {"base": {"fhir": row(refused_prior=2.0)}, "base+500k": {"fhir": empty}} + + winner, reasons = pick_winner(1, table, [], {"base": 0, "base+500k": 1}) + + assert winner == "base" + assert any("base+500k" in r and "no valid runs" in r for r in reasons) + + +def test_a_curly_apostrophe_still_grounds_the_fact() -> None: + page = "SHALL at least contain a contact's details or a reference to an organization" + + answer = f"must hold a contact{RIGHT_SINGLE_QUOTE}s details" + + assert is_grounded(answer, [fetch(content=page)], ("contact's details",)) + + +def test_a_control_that_used_the_web_before_failing_still_counts() -> None: + search = {"round": 1, "tool": "search", "target": "q", "result": "ok"} + summary = summarise([ + run(turn([]), index=0), + run(turn([search], error="ApolloError: overloaded"), index=1), + ], ()) + + assert summary["web_calls"] == 0 + assert summary["web_calls_max"] == 1 diff --git a/services/global_chat/tests/unit/test_web_tools_recording.py b/services/global_chat/tests/unit/test_web_tools_recording.py new file mode 100644 index 00000000..dafb0f07 --- /dev/null +++ b/services/global_chat/tests/unit/test_web_tools_recording.py @@ -0,0 +1,193 @@ +"""Unit tests for the web-tools recording planner, variants and config overrides.""" + +import pytest +from global_chat.config_loader import ConfigLoader +from global_chat.tests.web_tools import recording +from global_chat.tests.web_tools.recording import OverrideConfigLoader, RecordingPlanner, fingerprint +from global_chat.tests.web_tools.variants import ( + ENTRY_URLS, + EXAMPLE_URL, + FINDINGS, + SEARCH_FIRST, + STAGES, + resolve_variant, +) + + +def test_base_is_the_shipped_config() -> None: + variant = resolve_variant("base") + + assert (variant.web_search, variant.prompt_suffix, variant.inject_urls) == ({}, "", ()) + + +def test_variants_compose_left_to_right() -> None: + variant = resolve_variant("base+1c+findings+25k") + + assert variant.name == "base+1c+findings+25k" + assert SEARCH_FIRST in variant.prompt_suffix + assert FINDINGS in variant.prompt_suffix + assert variant.inject_urls == ENTRY_URLS + assert variant.web_search == {"max_content_tokens": 25000} + assert variant.complexity == len(["prompt", "code"]) + + +def test_an_unknown_variant_part_is_an_error() -> None: + with pytest.raises(KeyError): + resolve_variant("base+nope") + + +def test_every_stage_names_real_variants() -> None: + for stage in STAGES.values(): + for name in stage["variants"]: + resolve_variant(f"base+{name}") + + +OVERRIDE_TOKENS = 40000 + + +def test_override_loader_applies_the_variant_without_touching_the_files() -> None: + loader = OverrideConfigLoader(resolve_variant("base+findings+40k")) + fresh = ConfigLoader() + + assert loader.config["planner"]["web_search"]["max_content_tokens"] == OVERRIDE_TOKENS + assert loader.get_prompt("planner_web_tools_prompt").rstrip().endswith(FINDINGS) + assert fresh.config["planner"]["web_search"]["max_content_tokens"] != OVERRIDE_TOKENS + assert FINDINGS not in fresh.get_prompt("planner_web_tools_prompt") + + +def test_a_suffix_already_in_the_prompt_is_not_added_twice() -> None: + prompts = {"planner_web_tools_prompt": f"intro\n{SEARCH_FIRST}\n"} + + recording.apply_prompt_changes(prompts, resolve_variant("base+1a")) + + assert prompts["planner_web_tools_prompt"].count(SEARCH_FIRST) == 1 + + +def test_1d_drops_the_example_url_and_adds_the_search_first_line() -> None: + shipped = ConfigLoader().get_prompt("planner_web_tools_prompt") + prompt = OverrideConfigLoader(resolve_variant("base+1d")).get_prompt("planner_web_tools_prompt") + + assert EXAMPLE_URL in shipped + assert EXAMPLE_URL not in prompt + assert "https://hl7.org/fhir/R4/patient.html" not in prompt + assert "`https://hl7.org/fhir/R4/`" in prompt + assert prompt.rstrip().endswith(SEARCH_FIRST) + + +def test_a_removal_that_matches_nothing_is_an_error() -> None: + prompts = {"planner_web_tools_prompt": "no example here"} + + with pytest.raises(ValueError, match="not in"): + recording.apply_prompt_changes(prompts, resolve_variant("base+1d")) + + +def make_recording_planner(inject_urls: tuple[str, ...]) -> RecordingPlanner: + planner = RecordingPlanner.__new__(RecordingPlanner) + planner.inject_urls = inject_urls + planner._attachments = [] + planner.current_yaml = None + return planner + + +def test_injected_urls_are_appended_to_the_user_turn() -> None: + content = make_recording_planner(ENTRY_URLS)._build_user_content("What is Patient.gender?", None) + + assert content.startswith("What is Patient.gender?") + for url in ENTRY_URLS: + assert url in content + + +def test_nothing_is_appended_without_injected_urls() -> None: + assert make_recording_planner(())._build_user_content("q", None) == "q" + + +def test_the_fingerprint_changes_with_anything_that_changes_behaviour() -> None: + prints = {fingerprint(resolve_variant(name)) for name in ("base", "base+findings", "base+1b", "base+25k")} + + assert len(prints) == len(["base", "base+findings", "base+1b", "base+25k"]) + assert fingerprint(resolve_variant("base")) == fingerprint(resolve_variant("base")) + + +class Scenario: + turns = ("first", "second") + workflow_yaml = None + page = None + + +class FailingPlanner: + """Stands in for RecordingPlanner: one recorded round, then the API fails.""" + + def __init__(self, *_args: object, **_kwargs: object) -> None: + self.responses = [] + + def run(self, *_args: object, **_kwargs: object) -> None: + raise RuntimeError("overloaded") + + +def test_a_failed_turn_is_recorded_and_ends_the_scenario(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(recording, "RecordingPlanner", FailingPlanner) + + turns = recording.run_scenario(Scenario(), resolve_variant("base")) + + assert len(turns) == 1 + assert turns[0].error == "RuntimeError: overloaded" + assert turns[0].answer == "" + + +def test_the_shipped_prompt_is_the_measured_1a_prompt() -> None: + assert fingerprint(resolve_variant("base")) == fingerprint(resolve_variant("base+1a")) + + +def test_a_shipped_suffix_is_not_repeated_when_composed_with_another() -> None: + prompt = OverrideConfigLoader(resolve_variant("base+1a+findings")).get_prompt("planner_web_tools_prompt") + + assert prompt.count(SEARCH_FIRST) == 1 + assert FINDINGS in prompt + + +def test_the_fingerprint_covers_the_resolved_model_and_the_whole_planner_config( + monkeypatch: pytest.MonkeyPatch, +) -> None: + before = fingerprint(resolve_variant("base")) + + monkeypatch.setitem(recording.CLAUDE_MODELS, "claude-opus", "claude-opus-next") + after_model = fingerprint(resolve_variant("base")) + monkeypatch.undo() + + original_init = OverrideConfigLoader.__init__ + + def init_with_more_tool_calls(self: OverrideConfigLoader, variant: object) -> None: + original_init(self, variant) + self.config["planner"]["max_tool_calls"] = 99 + + monkeypatch.setattr(OverrideConfigLoader, "__init__", init_with_more_tool_calls) + after_config = fingerprint(resolve_variant("base")) + + assert before != after_model + assert before != after_config + + +def test_the_scenario_key_changes_when_the_scenario_content_changes() -> None: + class Edited: + turns = ("a different question",) + workflow_yaml = None + page = None + + assert recording.scenario_key(Scenario()) != recording.scenario_key(Edited()) + assert recording.scenario_key(Scenario()) == recording.scenario_key(Scenario()) + + +class UnbuildablePlanner: + """Stands in for RecordingPlanner when PlannerAgent itself refuses to build.""" + + def __init__(self, *_args: object, **_kwargs: object) -> None: + raise RuntimeError("ANTHROPIC_API_KEY not found") + + +def test_a_planner_that_cannot_be_built_is_recorded_not_raised(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(recording, "RecordingPlanner", UnbuildablePlanner) + + turns = recording.run_scenario(Scenario(), resolve_variant("base")) + + assert len(turns) == 1 + assert turns[0].error == "RuntimeError: ANTHROPIC_API_KEY not found" diff --git a/services/global_chat/tests/unit/test_web_tools_runner.py b/services/global_chat/tests/unit/test_web_tools_runner.py new file mode 100644 index 00000000..78d62c30 --- /dev/null +++ b/services/global_chat/tests/unit/test_web_tools_runner.py @@ -0,0 +1,45 @@ +"""Unit tests for the experiment runner's run cache.""" + +import json +from dataclasses import asdict +from pathlib import Path + +from global_chat.tests.web_tools import run_web_experiments as runner +from global_chat.tests.web_tools.metrics import RunRecord, TurnRecord +from global_chat.tests.web_tools.scenarios import SCENARIOS + +USAGE = {"input_tokens": 1, "output_tokens": 1, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0} +SCENARIO = SCENARIOS["fhir_shallow"] + + +def turn(error: str | None = None) -> TurnRecord: + return TurnRecord(answer="a", trace=[], usage=USAGE, seconds=1.0, rounds=1, error=error) + + +def write(path: Path, *turns: TurnRecord) -> None: + record = RunRecord(SCENARIO.id, "base+1a", 0, list(turns)) + path.write_text(json.dumps(asdict(record)), encoding="utf-8") + + +def test_variants_with_the_same_prompt_share_one_cache_file(tmp_path: Path) -> None: + """base and base+1a render the same prompt since 1a shipped.""" + assert runner.cache_path(tmp_path, "base", SCENARIO, 0) == runner.cache_path(tmp_path, "base+1a", SCENARIO, 0) + assert runner.cache_path(tmp_path, "base", SCENARIO, 0) != runner.cache_path(tmp_path, "base+1b", SCENARIO, 0) + + +def test_a_cached_successful_run_is_reused(tmp_path: Path) -> None: + path = runner.cache_path(tmp_path, "base", SCENARIO, 0) + write(path, turn()) + + assert runner.cached_run(path) is not None + + +def test_a_cached_failed_run_is_not_reused(tmp_path: Path) -> None: + path = runner.cache_path(tmp_path, "base", SCENARIO, 0) + write(path, turn(), turn(error="ApolloError: overloaded")) + + assert runner.cached_run(path) is None + + +def test_a_missing_cache_file_is_not_a_hit(tmp_path: Path) -> None: + assert runner.cached_run(tmp_path / "absent.json") is None diff --git a/services/global_chat/tests/unit/test_web_tools_scenarios.py b/services/global_chat/tests/unit/test_web_tools_scenarios.py new file mode 100644 index 00000000..e1012675 --- /dev/null +++ b/services/global_chat/tests/unit/test_web_tools_scenarios.py @@ -0,0 +1,39 @@ +"""Unit tests that the web-tools scenarios are well formed.""" + +from global_chat.tests.web_tools.scenarios import SCENARIOS +from global_chat.tests.web_tools.variants import STAGES + + +def test_every_stage_names_real_scenarios() -> None: + for stage in STAGES.values(): + for scenario_id in stage["scenarios"]: + assert scenario_id in SCENARIOS + + +def test_ids_match_their_keys() -> None: + assert all(key == scenario.id for key, scenario in SCENARIOS.items()) + + +def test_controls_have_no_facts_and_fhir_scenarios_do() -> None: + for scenario in SCENARIOS.values(): + if scenario.control: + assert scenario.facts == () + assert SCENARIOS["fhir_shallow"].facts + assert SCENARIOS["fhir_deep"].facts + + +def test_multi_turn_has_several_turns() -> None: + assert len(SCENARIOS["multi_turn"].turns) > 1 + assert all(len(s.turns) == 1 for k, s in SCENARIOS.items() if k != "multi_turn") + + +def test_the_code_edit_control_points_at_a_real_step() -> None: + scenario = SCENARIOS["control_code_edit"] + assert scenario.page is not None + assert scenario.workflow_yaml is not None + + assert scenario.page.rsplit("/", 1)[-1] in scenario.workflow_yaml + + +def test_fhir_deep_names_the_page_whose_truncation_it_measures() -> None: + assert SCENARIOS["fhir_deep"].source_page == "hl7.org/fhir/R4/patient.html" diff --git a/services/global_chat/tests/unit/test_web_tools_trace.py b/services/global_chat/tests/unit/test_web_tools_trace.py new file mode 100644 index 00000000..50f654b4 --- /dev/null +++ b/services/global_chat/tests/unit/test_web_tools_trace.py @@ -0,0 +1,177 @@ +"""Unit tests for the web-tools trace builder (test helper, not production code).""" + +from types import SimpleNamespace as Block + +import pytest +from global_chat.tests.web_tools.trace import build_trace, count, format_trace + +PATIENT = "https://hl7.org/fhir/R4/patient.html" + + +def use(block_id: str, name: str, **tool_input: str) -> Block: + return Block(type="server_tool_use", id=block_id, name=name, input=tool_input) + + +def result(block_id: str, block_type: str, content: object) -> Block: + return Block(type=block_type, tool_use_id=block_id, content=content) + + +def response(*blocks: Block) -> Block: + return Block(content=list(blocks)) + + +def fetched(text: str) -> Block: + return Block(type="web_fetch_result", url=PATIENT, content=Block(type="document", source=Block(type="text", data=text))) + + +def error(kind: str, code: str) -> dict: + return {"type": f"{kind}_tool_result_error", "error_code": code} + + +SEARCH_OK = [{"type": "web_search_result", "url": PATIENT}] + + +def test_calls_are_recorded_in_order_across_rounds() -> None: + trace = build_trace([ + response( + use("f1", "web_fetch", url=PATIENT), + result("f1", "web_fetch_tool_result", error("web_fetch", "url_not_in_prior_context")), + ), + response( + use("s1", "web_search", query="fhir patient"), + result("s1", "web_search_tool_result", SEARCH_OK), + use("f2", "web_fetch", url=PATIENT), + result("f2", "web_fetch_tool_result", fetched("Patient.gender 0..1")), + ), + ]) + + assert [(c["round"], c["tool"], c["result"]) for c in trace] == [ + (1, "fetch", "url_not_in_prior_context"), + (2, "search", "ok"), + (2, "fetch", "ok"), + ] + assert trace[0]["target"] == PATIENT + assert trace[1]["target"] == "fhir patient" + + +def test_a_successful_fetch_keeps_the_page_text() -> None: + trace = build_trace([response(use("f1", "web_fetch", url=PATIENT), result("f1", "web_fetch_tool_result", fetched("abc")))]) + + assert (trace[0]["content"], trace[0]["content_chars"]) == ("abc", len("abc")) + + +def test_results_are_paired_by_id_not_position() -> None: + trace = build_trace([response( + use("f1", "web_fetch", url="https://hl7.org/a"), + use("f2", "web_fetch", url="https://hl7.org/b"), + result("f2", "web_fetch_tool_result", error("web_fetch", "url_not_allowed")), + result("f1", "web_fetch_tool_result", fetched("page a")), + )]) + + assert [(c["target"], c["result"]) for c in trace] == [ + ("https://hl7.org/a", "ok"), + ("https://hl7.org/b", "url_not_allowed"), + ] + + +def test_a_result_can_arrive_in_a_later_round() -> None: + """pause_turn can split a call from its result across responses.""" + trace = build_trace([ + response(use("f1", "web_fetch", url=PATIENT)), + response(result("f1", "web_fetch_tool_result", fetched("late"))), + ]) + + assert trace[0]["round"] == 1 + assert trace[0]["result"] == "ok" + assert trace[0]["content"] == "late" + + +def test_a_call_with_no_result_is_marked_missing() -> None: + trace = build_trace([response(use("s1", "web_search", query="q"))]) + + assert trace[0]["result"] == "missing" + + +def test_a_non_text_fetch_has_no_content() -> None: + pdf = Block(type="web_fetch_result", url=PATIENT, content=Block(type="document", source=Block(type="base64", data="JVBERi0="))) + trace = build_trace([response(use("f1", "web_fetch", url=PATIENT), result("f1", "web_fetch_tool_result", pdf))]) + + assert trace[0]["result"] == "ok" + assert trace[0]["content"] is None + assert trace[0]["content_chars"] is None + + +def test_other_server_tools_keep_their_own_name() -> None: + trace = build_trace([response(use("c1", "code_execution", code="print(1)"))]) + + assert trace[0]["tool"] == "code_execution" + + +@pytest.mark.parametrize("code", ["url_not_allowed", "max_uses_exceeded", "url_not_accessible"]) +def test_each_error_code_is_kept(code: str) -> None: + trace = build_trace([response(use("f1", "web_fetch", url=PATIENT), result("f1", "web_fetch_tool_result", error("web_fetch", code)))]) + + assert trace[0]["result"] == code + assert "content" not in trace[0] + + +def test_dict_blocks_are_read_like_sdk_objects() -> None: + trace = build_trace([{"content": [ + {"type": "server_tool_use", "id": "s1", "name": "web_search", "input": {"query": "q"}}, + {"type": "web_search_tool_result", "tool_use_id": "s1", "content": SEARCH_OK}, + ]}]) + + assert trace == [{"round": 1, "tool": "search", "target": "q", "result": "ok"}] + + +def test_count_filters_by_tool_and_result() -> None: + trace = [ + {"round": 1, "tool": "fetch", "target": "a", "result": "url_not_in_prior_context"}, + {"round": 1, "tool": "search", "target": "q", "result": "ok"}, + {"round": 2, "tool": "fetch", "target": "a", "result": "ok"}, + ] + + assert [ + count(trace), + count(trace, tool="fetch"), + count(trace, result="ok"), + count(trace, tool="fetch", result="ok"), + ] == [3, 2, 2, 1] + + +def test_format_trace_lists_one_line_per_call() -> None: + refused, read = format_trace([ + {"round": 1, "tool": "fetch", "target": PATIENT, "result": "url_not_in_prior_context"}, + {"round": 2, "tool": "fetch", "target": PATIENT, "result": "ok", "content": "x" * 10, "content_chars": 10}, + ]).splitlines() + + assert "url_not_in_prior_context" in refused + assert PATIENT in read + assert "10 chars" in read + + +def test_format_trace_says_when_there_were_no_calls() -> None: + assert format_trace([]) == "(no web calls)" + + +def test_a_code_execution_result_is_paired_with_its_call() -> None: + trace = build_trace([response( + use("c1", "code_execution", code="print(page[:100])"), + result("c1", "code_execution_tool_result", Block(type="code_execution_result", stdout="ok")), + )]) + + assert (trace[0]["tool"], trace[0]["result"]) == ("code_execution", "ok") + + +def test_other_code_execution_results_are_paired_too() -> None: + trace = build_trace([response( + use("b1", "bash_code_execution", command="ls"), + result("b1", "bash_code_execution_tool_result", Block(type="bash_code_execution_result", stdout="")), + use("t1", "text_editor_code_execution", command="view"), + result("t1", "text_editor_code_execution_tool_result", Block(type="text_editor_code_execution_view_result")), + )]) + + assert [(c["tool"], c["result"]) for c in trace] == [ + ("bash_code_execution", "ok"), + ("text_editor_code_execution", "ok"), + ] diff --git a/services/global_chat/tests/web_tools/__init__.py b/services/global_chat/tests/web_tools/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/services/global_chat/tests/web_tools/calibrate.py b/services/global_chat/tests/web_tools/calibrate.py new file mode 100644 index 00000000..73848d19 --- /dev/null +++ b/services/global_chat/tests/web_tools/calibrate.py @@ -0,0 +1,90 @@ +"""Fetch long FHIR pages at several max_content_tokens and report where facts land. + +Answers two questions before any experiment money is spent: + 1. Does a web_fetch_20260209 result surface as a web_fetch_tool_result block + with the page text (what build_trace reads), or does dynamic filtering hide it? + 2. Which fact sits past the 10k truncation point but within reach at a larger size? + +Run from the repo root: + PYTHONUTF8=1 PYTHONPATH=services python -m poetry run python -m global_chat.tests.web_tools.calibrate +""" + +# ruff: noqa: T201 - a command-line report, where printing is its output + +import os +import sys +from pathlib import Path + +from dotenv import load_dotenv + +load_dotenv(Path(__file__).resolve().parents[3] / ".env") +load_dotenv() + +from anthropic import Anthropic # noqa: E402 +from global_chat.config_loader import ConfigLoader # noqa: E402 +from models import resolve_model # noqa: E402 + +from .trace import build_trace, format_trace # noqa: E402 + +PAGES = { + "https://hl7.org/fhir/R4/patient.html": [ + "male | female | other | unknown", + "replaced-by", + "seealso", + "pat-1", + "SHALL at least contain a contact's details or a reference to an organization", + "death-date", + "general-practitioner", + "link", + ], + "https://hl7.org/fhir/R4/observation.html": [ + "registered | preliminary | final | amended", + "obs-6", + "obs-7", + "code-value-concept", + "combo-code-value-quantity", + "component-value-concept", + ], +} +SIZES = (10000, 25000, 50000) + + +def main() -> None: + if os.getenv("ANTHROPIC_BASE_URL"): + sys.exit("ANTHROPIC_BASE_URL is set, the web tools need a direct api.anthropic.com key") + + model = resolve_model(ConfigLoader().config["planner"]["model"]) + client = Anthropic() + + for page, facts in PAGES.items(): + for size in SIZES: + tool = { + "type": "web_fetch_20260209", + "name": "web_fetch", + "max_uses": 1, + "max_content_tokens": size, + "allowed_domains": ["hl7.org"], + } + response = client.beta.messages.create( + model=model, + max_tokens=2048, + tools=[tool], + messages=[{"role": "user", "content": f"Fetch {page} and then reply with just the word done"}], + ) + print(f"\n=== {page} @ {size} tokens") + print("block types:", [getattr(b, "type", None) for b in response.content]) + trace = build_trace([response]) + print(format_trace(trace)) + text = next((c["content"] for c in trace if c.get("content")), None) + if text is None: + print("NO FETCHED TEXT VISIBLE TO build_trace") + continue + lower = " ".join(text.lower().split()) + for fact in facts: + at = lower.find(" ".join(fact.lower().split())) + where = f"char {at} of {len(lower)}" if at >= 0 else "ABSENT" + print(f" {fact!r}: {where}") + + +if __name__ == "__main__": + main() diff --git a/services/global_chat/tests/web_tools/metrics.py b/services/global_chat/tests/web_tools/metrics.py new file mode 100644 index 00000000..512a0485 --- /dev/null +++ b/services/global_chat/tests/web_tools/metrics.py @@ -0,0 +1,222 @@ +"""Records, metrics and winner selection for the web-tools experiment.""" + +from dataclasses import dataclass, field +from statistics import mean + +from .trace import count + +# Approximate list prices for the planner's model (claude-opus via models.py), +# in $ per token. Subagent calls not included. +PRICE_INPUT = 5 / 1_000_000 +PRICE_OUTPUT = 25 / 1_000_000 +PRICE_CACHE_WRITE = PRICE_INPUT * 1.25 +PRICE_CACHE_READ = PRICE_INPUT * 0.1 +# Web search bills per search. +PRICE_SEARCH = 10 / 1000 + +# Straight apostrophe for matching. +RIGHT_SINGLE_QUOTE = chr(0x2019) + +NUMERIC = ( + "web_calls", "refused_prior", "refused_allowlist", "followup_fetches", + "input_tokens", "seconds", "cost", +) + +# Stage -> (metric averaged over non-control scenarios, True when lower is better). +PRIMARY = {1: ("refused_prior", True), 2: ("grounded_rate", False), 3: ("followup_fetches", True)} + + +@dataclass +class TurnRecord: + answer: str + trace: list[dict] + usage: dict + seconds: float + rounds: int + downgraded: bool = False + error: str | None = None + + +@dataclass +class RunRecord: + scenario_id: str + variant: str + run_index: int + turns: list[TurnRecord] = field(default_factory=list) + + @classmethod + def from_dict(cls, data: dict) -> "RunRecord": + turns = [TurnRecord(**t) for t in data["turns"]] + return cls(data["scenario_id"], data["variant"], data["run_index"], turns) + + +def is_grounded(answer: str, trace: list[dict], facts: tuple[str, ...]) -> bool: + """Every fact is in the answer and in a fetched page, so it was read, not remembered.""" + return all(norm(fact) in norm(answer) for fact in facts) and facts_fetched(trace, facts) + + +def facts_fetched(trace: list[dict], facts: tuple[str, ...]) -> bool: + """Every fact appears in some fetched page: the truncation signal, independent of the answer.""" + pages = fetched_pages(trace) + return all(any(norm(fact) in page for page in pages) for fact in facts) + + +def run_metrics(run: RunRecord, facts: tuple[str, ...]) -> dict: + trace = [call for turn in run.turns for call in turn.trace] + answers = "\n".join(turn.answer for turn in run.turns) + return { + "web_calls": count(trace, tool="search") + count(trace, tool="fetch"), + "refused_prior": count(trace, result="url_not_in_prior_context"), + "refused_allowlist": count(trace, result="url_not_allowed"), + "followup_fetches": sum(count(turn.trace, tool="fetch") for turn in run.turns[1:]), + "grounded": is_grounded(answers, trace, facts) if facts else None, + "fact_fetched": facts_fetched(trace, facts) if facts else None, + "input_tokens": sum( + turn.usage.get(key, 0) + for turn in run.turns + for key in ("input_tokens", "cache_creation_input_tokens", "cache_read_input_tokens") + ), + "seconds": sum(turn.seconds for turn in run.turns), + "cost": sum(turn_cost(turn) for turn in run.turns), + } + + +def summarise(runs: list[RunRecord], facts: tuple[str, ...], source_page: str | None = None) -> dict: + """Mean metrics over the valid runs. Failed and downgraded runs are counted.""" + valid = [run for run in runs if not failed(run) and not downgraded(run)] + rows = [run_metrics(run, facts) for run in valid] + summary: dict = { + "n": len(runs), + "valid": len(valid), + "errors": sum(1 for run in runs if failed(run)), + "downgraded": sum(1 for run in runs if downgraded(run) and not failed(run)), + } + for key in NUMERIC: + summary[key] = mean(r[key] for r in rows) if rows else None + summary["web_calls_max"] = max(run_metrics(run, facts)["web_calls"] for run in runs) if runs else None + summary["grounded_rate"] = mean(r["grounded"] for r in rows) if rows and facts else None + summary["fact_fetched_rate"] = mean(r["fact_fetched"] for r in rows) if rows and facts else None + summary["source_page_rate"] = ( + mean(fetched_source(run, source_page) for run in valid) if valid and source_page else None + ) + return summary + + +def pick_winner( + stage: int, + table: dict[str, dict[str, dict]], + controls: list[str], + complexity: dict[str, int], +) -> tuple[str | None, list[str]]: + """Apply the spec's winner rules in order and say why at each step. + + 1. Disqualify a variant if a control run used the web, or if it has more + allowlist refusals than the carried-forward variant (the table's first key). + 2. Rank by the stage's primary metric over the non-control scenarios. + 3. Break ties by input tokens, then seconds, then the simpler change. + """ + reasons: list[str] = [] + baseline = next(iter(table)) + base_allowlist = sum(r.get("refused_allowlist") or 0 for r in table[baseline].values()) + metric, lower_is_better = PRIMARY[stage] + + eligible = [] + for variant, rows in table.items(): + if any(not s.get("valid") for s in rows.values()): + reasons.append(f"{variant}: disqualified, a scenario has no valid runs") + continue + if any((rows[s].get("web_calls_max") or 0) > 0 for s in controls if s in rows): + reasons.append(f"{variant}: disqualified, a control run used the web") + continue + if sum(r.get("refused_allowlist") or 0 for r in rows.values()) > base_allowlist: + reasons.append(f"{variant}: disqualified, more allowlist refusals than {baseline}") + continue + eligible.append(variant) + + if not eligible: + return None, reasons + + def rank(variant: str) -> tuple: + rows = table[variant] + scored = [s for s in rows if s not in controls] + primary = round(mean_over(rows, scored, metric), 2) + return ( + primary if lower_is_better else -primary, + round(mean_over(rows, scored, "input_tokens")), + round(mean_over(rows, scored, "seconds"), 1), + complexity.get(variant, 0), + ) + + for variant in eligible: + primary, tokens, seconds, simplicity = rank(variant) + reasons.append( + f"{variant}: {metric}={abs(primary)} input_tokens={tokens} seconds={seconds} complexity={simplicity}", + ) + return min(eligible, key=rank), reasons + + +def format_table(table: dict[str, dict[str, dict]]) -> str: + header = ( + "| variant | scenario | valid/n | web calls | refused prior | refused allowlist " + "| grounded | fact fetched | source page | follow-up fetches | input tok | sec | $ |\n" + "|---|---|---|---|---|---|---|---|---|---|---|---|---|" + ) + lines = [header] + for variant, rows in table.items(): + for scenario_id, s in rows.items(): + lines.append( + f"| {variant} | {scenario_id} | {s['valid']}/{s['n']} | {cell(s['web_calls'])} " + f"| {cell(s['refused_prior'])} | {cell(s['refused_allowlist'])} " + f"| {cell(s['grounded_rate'])} | {cell(s['fact_fetched_rate'])} | {cell(s['source_page_rate'])} " + f"| {cell(s['followup_fetches'])} | {cell(s['input_tokens'], 0)} " + f"| {cell(s['seconds'], 1)} | {cell(s['cost'], 3)} |", + ) + return "\n".join(lines) + + +def norm(text: str) -> str: + return " ".join(text.replace(RIGHT_SINGLE_QUOTE, "'").lower().split()) + + +def fetched_pages(trace: list[dict]) -> list[str]: + return [norm(call["content"]) for call in trace if call.get("content")] + + +def fetched_source(run: RunRecord, source_page: str) -> bool: + return any( + call["tool"] == "fetch" and call["result"] == "ok" and source_page in call["target"] + for turn in run.turns + for call in turn.trace + ) + + +def turn_cost(turn: TurnRecord) -> float: + usage = turn.usage + return ( + usage.get("input_tokens", 0) * PRICE_INPUT + + usage.get("output_tokens", 0) * PRICE_OUTPUT + + usage.get("cache_creation_input_tokens", 0) * PRICE_CACHE_WRITE + + usage.get("cache_read_input_tokens", 0) * PRICE_CACHE_READ + + count(turn.trace, tool="search") * PRICE_SEARCH + ) + + +def failed(run: RunRecord) -> bool: + return any(turn.error for turn in run.turns) + + +def downgraded(run: RunRecord) -> bool: + return any(turn.downgraded for turn in run.turns) + + +def mean_over(rows: dict[str, dict], scenarios: list[str], key: str) -> float: + values = [rows[s][key] for s in scenarios if rows[s].get(key) is not None] + return mean(values) if values else 0.0 + + +def cell(value: object, digits: int = 2) -> str: + if value is None: + return "-" + if isinstance(value, float): + return f"{value:.{digits}f}" + return str(value) diff --git a/services/global_chat/tests/web_tools/recording.py b/services/global_chat/tests/web_tools/recording.py new file mode 100644 index 00000000..f2653d98 --- /dev/null +++ b/services/global_chat/tests/web_tools/recording.py @@ -0,0 +1,137 @@ +"""Run the planner in-process and keep every raw API response it received.""" + +import hashlib +import json +import time +from typing import Protocol + +from global_chat.config_loader import ConfigLoader +from global_chat.planner import PlannerAgent +from models import CLAUDE_MODELS + +from .metrics import TurnRecord +from .trace import build_trace +from .variants import INJECT_TEMPLATE, Variant + +WEB_PROMPT_KEY = "planner_web_tools_prompt" +USAGE_FIELDS = ("input_tokens", "output_tokens", "cache_creation_input_tokens", "cache_read_input_tokens") + + +class PlayableScenario(Protocol): + """What run_scenario reads from a scenario; scenarios.Scenario satisfies it.""" + + @property + def turns(self) -> tuple[str, ...]: ... + + @property + def workflow_yaml(self) -> str | None: ... + + @property + def page(self) -> str | None: ... + + +class OverrideConfigLoader(ConfigLoader): + """The shipped config and prompts with one variant applied in memory.""" + + def __init__(self, variant: Variant) -> None: + super().__init__() + self.config["planner"]["web_search"].update(variant.web_search) + apply_prompt_changes(self.prompts["prompts"], variant) + + +class RecordingPlanner(PlannerAgent): + """PlannerAgent with web search on, recording each round's raw response.""" + + def __init__(self, config_loader: ConfigLoader, inject_urls: tuple[str, ...] = (), api_key: str | None = None) -> None: + super().__init__(config_loader, api_key=api_key, web_search=True) + self.inject_urls = tuple(inject_urls) + self.responses: list = [] + + def _call_api(self, *args: object, **kwargs: object) -> object: + response = super()._call_api(*args, **kwargs) + self.responses.append(response) + return response + + def _build_user_content(self, content: str, page: str | None) -> str: + user_content = super()._build_user_content(content, page) + if self.inject_urls: + user_content += "\n\n" + INJECT_TEMPLATE.format(urls=", ".join(self.inject_urls)) + return user_content + + +def apply_prompt_changes(prompts: dict, variant: Variant) -> None: + """Remove the variant's exact text from the web prompt, then append its suffix. + + A removal that matches nothing raises, so a prompts.yaml edit can never turn + a variant silently into a different one. + """ + text = prompts[WEB_PROMPT_KEY] + for removal in variant.prompt_removals: + if removal not in text: + raise ValueError(f"{removal!r} not in {WEB_PROMPT_KEY}") + text = text.replace(removal, "") + for line in variant.prompt_suffix.split("\n") if variant.prompt_suffix else []: + if line not in text: + text = text.rstrip("\n") + "\n" + line + "\n" + prompts[WEB_PROMPT_KEY] = text + + +def fingerprint(variant: Variant) -> str: + """A short hash of everything that changes planner behaviour, used to key cached runs.""" + loader = OverrideConfigLoader(variant) + planner_config = loader.config["planner"] + material = { + "model": CLAUDE_MODELS.get(planner_config.get("model"), planner_config.get("model")), + "planner": planner_config, + "system": loader.get_prompt("planner_system_prompt"), + "web_prompt": loader.get_prompt(WEB_PROMPT_KEY), + "inject": list(variant.inject_urls), + } + return hashlib.sha256(json.dumps(material, sort_keys=True).encode()).hexdigest()[:10] + + +def scenario_key(scenario: PlayableScenario) -> str: + """A short hash of what a scenario sends, so editing its turns never reuses old answers.""" + material = [list(scenario.turns), scenario.workflow_yaml, scenario.page] + return hashlib.sha256(json.dumps(material).encode()).hexdigest()[:6] + + +def run_scenario(scenario: PlayableScenario, variant: Variant) -> list[TurnRecord]: + """Play every turn of a scenario live, carrying history the way return_history does.""" + loader = OverrideConfigLoader(variant) + history: list[dict] = [] + turns: list[TurnRecord] = [] + + for content in scenario.turns: + responses: list = [] + start = time.monotonic() + try: + planner = RecordingPlanner(loader, inject_urls=variant.inject_urls) + responses = planner.responses + result = planner.run(content, scenario.workflow_yaml, scenario.page, history, stream=False) + except Exception as error: # recorded as a failed run, including a planner that cannot be built + turns.append(TurnRecord( + answer="", + trace=build_trace(responses), + usage=usage(responses), + seconds=time.monotonic() - start, + rounds=len(responses), + error=f"{type(error).__name__}: {error}", + )) + break + + turns.append(TurnRecord( + answer=result.response, + trace=build_trace(planner.responses), + usage=usage(planner.responses), + seconds=time.monotonic() - start, + rounds=len(planner.responses), + downgraded=bool(result.meta.get("web_search_downgraded")), + )) + history = result.history + + return turns + + +def usage(responses: list) -> dict: + return {key: sum(getattr(r.usage, key, 0) or 0 for r in responses) for key in USAGE_FIELDS} diff --git a/services/global_chat/tests/web_tools/run_web_experiments.py b/services/global_chat/tests/web_tools/run_web_experiments.py new file mode 100644 index 00000000..9fb87ed4 --- /dev/null +++ b/services/global_chat/tests/web_tools/run_web_experiments.py @@ -0,0 +1,119 @@ +"""Run the staged web-tools experiment and print one comparison table per stage. + +Runs are cached in tmp/ keyed by variant, config fingerprint, scenario and run +index. Delete a file to re-run that cell. + +From the repo root: + PYTHONUTF8=1 PYTHONPATH=services python -m poetry run python -m global_chat.tests.web_tools.run_web_experiments --stage 0 + ... --stage 1 --carry base + ... --stage 2 --carry base+1c --variants 20k,30k + ... --scenario fhir_deep --carry base --runs 1 +""" + +# ruff: noqa: T201 - a command-line report, where printing is its output + +import argparse +import json +import os +import sys +from dataclasses import asdict +from pathlib import Path + +from dotenv import load_dotenv + +load_dotenv(Path(__file__).resolve().parents[3] / ".env") +load_dotenv() + +from .metrics import RunRecord, format_table, pick_winner, summarise # noqa: E402 +from .recording import fingerprint, run_scenario, scenario_key # noqa: E402 +from .scenarios import SCENARIOS, Scenario # noqa: E402 +from .trace import format_trace # noqa: E402 +from .variants import STAGES, resolve_variant # noqa: E402 + +TMP = Path(__file__).parent / "tmp" + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument("--stage", type=int, choices=sorted(STAGES), default=0) + parser.add_argument("--carry", default="base", help="variant carried forward from the previous stage") + parser.add_argument("--variants", default="", help="comma-separated parts overriding the stage's variants") + parser.add_argument("--scenario", default="", help="run one scenario for the carried variant and print it in full") + parser.add_argument("--runs", type=int, default=3) + args = parser.parse_args() + + if os.getenv("ANTHROPIC_BASE_URL"): + sys.exit("ANTHROPIC_BASE_URL is set, the web tools need a direct api.anthropic.com key") + + if args.scenario: + contestants, scenario_ids = [args.carry], [args.scenario] + else: + stage = STAGES[args.stage] + parts = [p for p in args.variants.split(",") if p] or stage["variants"] + contestants = [args.carry] + [f"{args.carry}+{p}" for p in parts] + scenario_ids = stage["scenarios"] + + table: dict[str, dict[str, dict]] = {} + for variant_name in contestants: + table[variant_name] = {} + for scenario_id in scenario_ids: + scenario = SCENARIOS[scenario_id] + runs = [load_or_run(variant_name, scenario, i) for i in range(args.runs)] + table[variant_name][scenario_id] = summarise(runs, scenario.facts, scenario.source_page) + if args.scenario: + for run in runs: + print_run(run) + + report = format_table(table) + if not args.scenario and args.stage > 0: + controls = [s for s in scenario_ids if SCENARIOS[s].control] + complexity = {v: resolve_variant(v).complexity for v in table} + winner, reasons = pick_winner(args.stage, table, controls, complexity) + report += "\n\n" + "\n".join(reasons) + f"\n\nwinner: {winner}" + + print("\n" + report) + if not args.scenario: + TMP.mkdir(exist_ok=True) + (TMP / f"stage-{args.stage}__{args.carry}.md").write_text(report + "\n", encoding="utf-8") + + +def load_or_run(variant_name: str, scenario: Scenario, run_index: int) -> RunRecord: + path = cache_path(TMP, variant_name, scenario, run_index) + cached = cached_run(path) + if cached is not None: + return RunRecord(cached.scenario_id, variant_name, run_index, cached.turns) + + print(f"running {variant_name} / {scenario.id} / run {run_index}", flush=True) + record = RunRecord(scenario.id, variant_name, run_index, run_scenario(scenario, resolve_variant(variant_name))) + TMP.mkdir(exist_ok=True) + path.write_text(json.dumps(asdict(record), indent=2), encoding="utf-8") + return record + + +def cache_path(directory: Path, variant_name: str, scenario: Scenario, run_index: int) -> Path: + key = fingerprint(resolve_variant(variant_name)) + return directory / f"{key}__{scenario.id}-{scenario_key(scenario)}__run-{run_index}.json" + + +def cached_run(path: Path) -> RunRecord | None: + """A cached run, or None to re-run it. A failed run is retried.""" + if not path.exists(): + return None + record = RunRecord.from_dict(json.loads(path.read_text(encoding="utf-8"))) + if any(turn.error for turn in record.turns): + print(f"retrying {path.name}, it failed last time", flush=True) + return None + return record + + +def print_run(run: RunRecord) -> None: + for number, turn in enumerate(run.turns, start=1): + print(f"\n--- {run.variant} / {run.scenario_id} / run {run.run_index} / turn {number}") + if turn.error: + print(f"ERROR: {turn.error}") + print(format_trace(turn.trace)) + print(f"\n{turn.answer}") + + +if __name__ == "__main__": + main() diff --git a/services/global_chat/tests/web_tools/scenarios.py b/services/global_chat/tests/web_tools/scenarios.py new file mode 100644 index 00000000..cac4411d --- /dev/null +++ b/services/global_chat/tests/web_tools/scenarios.py @@ -0,0 +1,98 @@ +"""The six web-tools scenarios, shared by the runner and the pass/fail suite. + +Each scenario has a trap: a reason the model might use the web wrongly. +""" + +from dataclasses import dataclass + + +@dataclass(frozen=True) +class Scenario: + id: str + turns: tuple[str, ...] + workflow_yaml: str | None = None + page: str | None = None + # Ground truth, each string must be in the answer and in a fetched page. + facts: tuple[str, ...] = () + # The control must make no web calls. + control: bool = False + source_page: str | None = None + + +FHIR_MAPPING_YAML = """\ +name: commcare-to-fhir +jobs: + map-patient: + id: job-map-patient + name: Map Patient + adaptor: "@openfn/language-common@latest" + body: | + fn(state => { + const patient = { + resourceType: 'Patient', + firstName: state.data.first_name, + lastName: state.data.last_name, + gender: state.data.sex, + }; + return { ...state, patient }; + }); +triggers: + webhook: + id: trigger-webhook + type: webhook + enabled: true +edges: + webhook->map-patient: + id: edge-webhook-map + source_trigger: webhook + target_job: map-patient + condition_type: always + enabled: true +""" + +SCENARIOS = {s.id: s for s in ( + # docs.openfn.org is allowlisted, which tempts a fetch; search_documentation is the right tool. + Scenario( + id="control_openfn_concept", + turns=("What does each() do in OpenFn job code, and when should I use it instead of fn()?",), + control=True, + ), + # FHIR vocabulary is found everywhere, but this is a pure JavaScript edit. + Scenario( + id="control_code_edit", + turns=("In this step, rename the firstName field to given and wrap its value in an array.",), + workflow_yaml=FHIR_MAPPING_YAML, + page="workflows/commcare-to-fhir/map-patient", + control=True, + ), + # Positive case: the fact is inside what a 10k fetch returns. + # Calibrated, where the pat-1 constraint text sits at char ~24.5k of the ~26k a 10k fetch returns. + Scenario( + id="fhir_shallow", + turns=("In FHIR R4, what rule applies to each entry in Patient.contact? What must it contain at minimum?",), + facts=("contact's details", "reference to an organization"), + ), + # Calibrated: absent from a 10k fetch of patient.html, present from 25k (char ~34.9k). + # The codes also sit near the top of the short valueset-link-type.html page, so a + # grounded answer does not by itself show truncation was avoided. + Scenario( + id="fhir_deep", + turns=("In FHIR R4, what codes can Patient.link.type take, and what does each mean?",), + facts=("replaced-by", "seealso"), + source_page="hl7.org/fhir/R4/patient.html", + ), + # The model knows docs.dhis2.org from training, but it is not on the allowlist. + Scenario( + id="off_allowlist", + turns=("What fields does the DHIS2 tracker API (/api/tracker) accept when creating an event?",), + ), + # Turn 2 is answerable from turn 1's page, and turn 3 barely needs it. + Scenario( + id="multi_turn", + turns=( + "Which fields does a FHIR R4 Patient have for contact information?", + "Which of those fields can repeat?", + "How would I map a CommCare phone number into it?", + ), + ), +)} diff --git a/services/global_chat/tests/web_tools/trace.py b/services/global_chat/tests/web_tools/trace.py new file mode 100644 index 00000000..e5c7445d --- /dev/null +++ b/services/global_chat/tests/web_tools/trace.py @@ -0,0 +1,93 @@ +"""Turn the planner's raw API responses into an ordered record of web tool calls. + +Test helper only. +""" + +from typing import Any + +WEB_TOOL_NAMES = {"web_search": "search", "web_fetch": "fetch"} +RESULT_BLOCK_SUFFIX = "_tool_result" + + +def build_trace(responses: list) -> list[dict]: + """One entry per server tool call, in the order the model made them. + + Results are paired to calls by tool_use_id across all responses, because a + pause_turn can split a call from its result. + """ + calls: list[dict] = [] + by_id: dict[object, dict] = {} + + for round_index, response in enumerate(responses, start=1): + for block in field(response, "content") or []: + block_type = field(block, "type") + if block_type == "server_tool_use": + name = str(field(block, "name")) + tool_input = field(block, "input") or {} + entry = { + "round": round_index, + "tool": WEB_TOOL_NAMES.get(name, name), + "target": tool_input.get("url") or tool_input.get("query") or "", + "result": "missing", + } + calls.append(entry) + by_id[field(block, "id")] = entry + elif str(block_type).endswith(RESULT_BLOCK_SUFFIX): + entry = by_id.get(field(block, "tool_use_id")) + if entry is None: + continue + content = field(block, "content") + entry["result"] = result_code(content) + if entry["tool"] == "fetch" and entry["result"] == "ok": + text = fetched_text(content) + entry["content"] = text + entry["content_chars"] = len(text) if text is not None else None + + return calls + + +def count(trace: list[dict], tool: str | None = None, result: str | None = None) -> int: + """How many calls match the given tool and/or result.""" + return sum( + 1 + for call in trace + if (tool is None or call["tool"] == tool) and (result is None or call["result"] == result) + ) + + +def format_trace(trace: list[dict]) -> str: + """One readable line per call, for failure messages and the runner's output.""" + if not trace: + return "(no web calls)" + lines = [] + for call in trace: + size = f" {call['content_chars']} chars" if call.get("content_chars") is not None else "" + lines.append(f"r{call['round']} {call['tool']:<7} {call['result']:<26} {call['target']}{size}") + return "\n".join(lines) + + +def field(obj: object, name: str) -> Any: # noqa: ANN401 + """Read a field from an SDK block or a plain dict; the unit fakes use both.""" + if isinstance(obj, dict): + return obj.get(name) + return getattr(obj, name, None) + + +def result_code(content: object) -> str: + """'ok', or the error_code of a failed result. A list is always a search success.""" + if content is None: + return "missing" + if isinstance(content, list): + return "ok" + if str(field(content, "type") or "").endswith("_error"): + return str(field(content, "error_code") or "unknown_error") + return "ok" + + +def fetched_text(content: object) -> str | None: + """The plain text of a successful fetch, or None when it is not text (e.g. a PDF).""" + source = field(field(content, "content"), "source") + if field(source, "type") != "text": + return None + data = field(source, "data") + return data if isinstance(data, str) else None diff --git a/services/global_chat/tests/web_tools/variants.py b/services/global_chat/tests/web_tools/variants.py new file mode 100644 index 00000000..22bf2f5c --- /dev/null +++ b/services/global_chat/tests/web_tools/variants.py @@ -0,0 +1,88 @@ +"""The experiment's variants and stages. A variant name composes parts with '+'.""" + +import re +from dataclasses import dataclass, field + +SEARCH_FIRST = ( + "- `web_fetch` can only open a URL that already appears in this conversation: in the " + "user's message, or in an earlier search or fetch result. URLs you remember, and URLs in " + "these instructions, are refused. Search first, then fetch a URL from the results." +) + +FINDINGS = ( + "- Fetched pages are not kept after this turn. When you used them, state the specific " + "facts you relied on and name the page, so a follow-up question can build on your " + "answer without fetching again." +) + +EXAMPLE_URL = " (e.g. `https://hl7.org/fhir/R4/patient.html`)" + +# One entry page per allowed domain in config.yaml. +ENTRY_URLS = ("https://hl7.org/fhir/R4/", "https://docs.openfn.org/") + +INJECT_TEMPLATE = "(Pages you can open with web_fetch: {urls})" + + +@dataclass +class Variant: + name: str + web_search: dict = field(default_factory=dict) + prompt_suffix: str = "" + # Text removed from the web prompt before the suffix is appended. + prompt_removals: tuple[str, ...] = () + inject_urls: tuple[str, ...] = () + # 0 = no change, 1 = prompt or config change, 2 = production code change. + complexity: int = 0 + + +PARTS = { + "base": Variant("base"), + "1a": Variant("1a", prompt_suffix=SEARCH_FIRST, complexity=1), + "1b": Variant("1b", inject_urls=ENTRY_URLS, complexity=2), + "1c": Variant("1c", prompt_suffix=SEARCH_FIRST, inject_urls=ENTRY_URLS, complexity=2), + "1d": Variant("1d", prompt_suffix=SEARCH_FIRST, prompt_removals=(EXAMPLE_URL,), complexity=1), + "findings": Variant("findings", prompt_suffix=FINDINGS, complexity=1), +} + +STAGES = { + 0: {"variants": [], "scenarios": [ + "control_openfn_concept", "control_code_edit", "fhir_shallow", + "fhir_deep", "off_allowlist", "multi_turn", + ]}, + 1: {"variants": ["1a", "1b", "1c"], "scenarios": [ + "control_openfn_concept", "control_code_edit", "fhir_shallow", "fhir_deep", + ]}, + 2: {"variants": ["25k", "50k"], "scenarios": ["fhir_shallow", "fhir_deep"]}, + 3: {"variants": ["findings"], "scenarios": ["multi_turn"]}, +} + + +def resolve_variant(name: str) -> Variant: + """Compose 'base+1c+25k' into one Variant; later parts win on conflicting config keys.""" + web_search: dict = {} + suffixes: list[str] = [] + removals: list[str] = [] + inject: tuple[str, ...] = () + complexity = 0 + for part in (parse_part(p) for p in name.split("+")): + web_search.update(part.web_search) + if part.prompt_suffix and part.prompt_suffix not in suffixes: + suffixes.append(part.prompt_suffix) + removals.extend(r for r in part.prompt_removals if r not in removals) + inject = inject or part.inject_urls + complexity = max(complexity, part.complexity) + return Variant( + name, + web_search=web_search, + prompt_suffix="\n".join(suffixes), + prompt_removals=tuple(removals), + inject_urls=inject, + complexity=complexity, + ) + + +def parse_part(name: str) -> Variant: + size = re.fullmatch(r"(\d+)k", name) + if size: + return Variant(name, web_search={"max_content_tokens": int(size.group(1)) * 1000}, complexity=1) + return PARTS[name]