Skip to content

Commit 004cee8

Browse files
committed
unit tests for api tool calling
1 parent 99f98b8 commit 004cee8

16 files changed

Lines changed: 903 additions & 146 deletions

tests/test_agentic_loop.py

Lines changed: 135 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,9 +5,9 @@
55

66
import pytest
77

8-
from src.tool_classifier.agentic_loop import AgenticLoop
9-
from src.tool_classifier.enums import AgenticLoopStatus
10-
from src.tool_classifier.param_extractor import ParamExtractionResult
8+
from tool_classifier.agentic_loop import AgenticLoop
9+
from tool_classifier.enums import AgenticLoopStatus
10+
from tool_classifier.param_extractor import ParamExtractionResult
1111

1212

1313
# ---------------------------------------------------------------------------
@@ -540,6 +540,7 @@ async def fake_to_thread(fn: Any, *args: Any, **kwargs: Any) -> Any:
540540
_HISTORY,
541541
{"validFrom": "2026-01-01"},
542542
"en",
543+
1,
543544
)
544545

545546

@@ -963,3 +964,134 @@ async def test_user_exit_during_stream_returns_empty_tokens(self) -> None:
963964
assert tokens == []
964965
# Collected params returned unchanged on exit
965966
assert result.collected_params == {"validFrom": "2026-01-01"}
967+
968+
969+
# ---------------------------------------------------------------------------
970+
# seeded_params — L2 param_update pre-population at turn 0
971+
# ---------------------------------------------------------------------------
972+
973+
974+
class TestSeededParamsTurn0:
975+
"""Verify that seeded_params from L2 follow-up routing are merged into
976+
collected_params at turn 0 only, with collected_params taking priority."""
977+
978+
@pytest.mark.asyncio
979+
async def test_seeded_params_merged_at_turn_0(self) -> None:
980+
"""seeded_params are prepended to collected_params when turn_count=0."""
981+
# Extractor returns only validFrom as newly extracted; countryIsoCode comes
982+
# from seeded_params.
983+
extractor_mock = _make_extractor_mock(
984+
_extraction(
985+
{"validFrom": "2026-01-01"},
986+
[], # nothing missing — both params will be present after seed merge
987+
"none",
988+
)
989+
)
990+
loop = _make_loop(extractor_mock)
991+
992+
result = await loop.run_turn(
993+
chat_id=_CHAT_ID,
994+
user_message="January 2026",
995+
conversation_history=[],
996+
params_schema=_SCHEMA_TWO_REQUIRED,
997+
collected_params={},
998+
turn_count=0,
999+
seeded_params={"countryIsoCode": "EE"},
1000+
)
1001+
1002+
# Both params present → COMPLETED
1003+
assert result.status == AgenticLoopStatus.COMPLETED
1004+
assert result.collected_params.get("countryIsoCode") == "EE"
1005+
assert result.collected_params.get("validFrom") == "2026-01-01"
1006+
1007+
@pytest.mark.asyncio
1008+
async def test_collected_params_override_seeded_params(self) -> None:
1009+
"""collected_params values beat seeded_params when the key overlaps."""
1010+
extractor_mock = _make_extractor_mock(
1011+
_extraction(
1012+
{"validFrom": "2026-06-01"},
1013+
[],
1014+
"none",
1015+
)
1016+
)
1017+
loop = _make_loop(extractor_mock)
1018+
1019+
result = await loop.run_turn(
1020+
chat_id=_CHAT_ID,
1021+
user_message="June 2026",
1022+
conversation_history=[],
1023+
params_schema=_SCHEMA_TWO_REQUIRED,
1024+
collected_params={"countryIsoCode": "LV"}, # explicit value takes priority
1025+
turn_count=0,
1026+
seeded_params={"countryIsoCode": "EE"}, # seeded value must be overridden
1027+
)
1028+
1029+
assert result.collected_params.get("countryIsoCode") == "LV"
1030+
1031+
@pytest.mark.asyncio
1032+
async def test_seeded_params_not_applied_on_subsequent_turns(self) -> None:
1033+
"""seeded_params are ignored when turn_count > 0."""
1034+
extractor_mock = _make_extractor_mock(
1035+
_extraction({}, ["countryIsoCode", "validFrom"], "Which country and date?")
1036+
)
1037+
loop = _make_loop(extractor_mock)
1038+
1039+
result = await loop.run_turn(
1040+
chat_id=_CHAT_ID,
1041+
user_message="hello",
1042+
conversation_history=[],
1043+
params_schema=_SCHEMA_TWO_REQUIRED,
1044+
collected_params={},
1045+
turn_count=1, # NOT turn 0 → seeded_params must be ignored
1046+
seeded_params={"countryIsoCode": "EE", "validFrom": "2026-01-01"},
1047+
)
1048+
1049+
# Even though seeded_params would satisfy all required params, they should
1050+
# not be applied after turn 0 → still NEEDS_INPUT
1051+
assert result.status == AgenticLoopStatus.NEEDS_INPUT
1052+
# seeded values not present in collected_params
1053+
assert "countryIsoCode" not in result.collected_params
1054+
assert "validFrom" not in result.collected_params
1055+
1056+
@pytest.mark.asyncio
1057+
async def test_seeded_params_none_does_not_raise(self) -> None:
1058+
"""Passing seeded_params=None (default) at turn 0 behaves normally."""
1059+
extractor_mock = _make_extractor_mock(
1060+
_extraction({}, ["countryIsoCode", "validFrom"], "Which country?")
1061+
)
1062+
loop = _make_loop(extractor_mock)
1063+
1064+
result = await loop.run_turn(
1065+
chat_id=_CHAT_ID,
1066+
user_message="hello",
1067+
conversation_history=[],
1068+
params_schema=_SCHEMA_TWO_REQUIRED,
1069+
collected_params={},
1070+
turn_count=0,
1071+
seeded_params=None,
1072+
)
1073+
1074+
assert result.status == AgenticLoopStatus.NEEDS_INPUT
1075+
1076+
@pytest.mark.asyncio
1077+
async def test_seeded_params_partial_fill_still_asks_for_missing(self) -> None:
1078+
"""seeded_params satisfy only one of two required params → still NEEDS_INPUT."""
1079+
extractor_mock = _make_extractor_mock(
1080+
_extraction({}, ["validFrom"], "From which date?")
1081+
)
1082+
loop = _make_loop(extractor_mock)
1083+
1084+
result = await loop.run_turn(
1085+
chat_id=_CHAT_ID,
1086+
user_message="Estonia",
1087+
conversation_history=[],
1088+
params_schema=_SCHEMA_TWO_REQUIRED,
1089+
collected_params={},
1090+
turn_count=0,
1091+
seeded_params={"countryIsoCode": "EE"}, # only one param seeded
1092+
)
1093+
1094+
# validFrom still missing → NEEDS_INPUT
1095+
assert result.status == AgenticLoopStatus.NEEDS_INPUT
1096+
# But seeded countryIsoCode should be present
1097+
assert result.collected_params.get("countryIsoCode") == "EE"

tests/test_api_caller.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -7,16 +7,16 @@
77
import httpx
88
import pytest
99

10-
from src.tool_classifier.api_caller import APICaller, CircuitBreaker
11-
from src.tool_classifier.constants import (
10+
from tool_classifier.api_caller import APICaller, CircuitBreaker
11+
from tool_classifier.constants import (
1212
CB_STATE_CLOSED,
1313
CB_STATE_HALF_OPEN,
1414
CB_STATE_OPEN,
1515
CIRCUIT_BREAKER_OPEN_MESSAGES,
1616
SERVICE_TIMEOUT_MESSAGES,
1717
SERVICE_UNAVAILABLE_MESSAGES,
1818
)
19-
from src.tool_classifier.models import APICallResult
19+
from tool_classifier.models import APICallResult
2020

2121

2222
# ---------------------------------------------------------------------------

tests/test_api_response_formatter.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@
99
import dspy.streaming
1010
import pytest
1111

12-
from src.tool_classifier.api_response_formatter import (
12+
from tool_classifier.api_response_formatter import (
1313
APIResponseFormatterModule,
1414
_FORMATTER_ERROR_MESSAGES,
1515
)

tests/test_api_semantic_searcher.py

Lines changed: 58 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -448,13 +448,6 @@ async def test_multiple_medium_triggers_disambiguation(self) -> None:
448448
client.get = AsyncMock(return_value=count_resp)
449449
client.post = AsyncMock(side_effect=[dense_resp, hybrid_resp])
450450

451-
mock_disambiguator = MagicMock()
452-
mock_disambiguator.return_value = None # "forward" returns string or None
453-
# Wrap in a module-like object that has a forward() callable via __call__
454-
mock_disambiguator_module = MagicMock()
455-
mock_disambiguator_module.forward = MagicMock(return_value="ep-holidays")
456-
mock_disambiguator_module.__call__ = MagicMock(return_value="ep-holidays")
457-
458451
# Inject our disambiguator — searcher calls self._disambiguator(query, candidates)
459452
# which in turn calls forward() via __call__
460453
async_disambiguator = MagicMock()
@@ -507,6 +500,64 @@ async def test_disambiguation_rejects_all_returns_empty(self) -> None:
507500

508501
assert results == []
509502

503+
@pytest.mark.asyncio
504+
async def test_disambiguation_rejects_all_multi_candidates_returns_top_with_hint(
505+
self,
506+
) -> None:
507+
"""Disambiguator returns None for >1 medium candidates → top candidate returned
508+
with multi_intent_hint=True and llm_validated=False so IntentDecomposer gate
509+
can run in the classifier."""
510+
cos_a = API_TOOL_MIN_THRESHOLD + 0.08 # higher cosine → becomes 'top'
511+
cos_b = API_TOOL_MIN_THRESHOLD + 0.02
512+
513+
dense_points = [
514+
_point({**_EP_HOLIDAYS}, cos_a),
515+
_point({**_EP_WEATHER}, cos_b),
516+
]
517+
hybrid_points = [
518+
_point({**_EP_HOLIDAYS}, 0.012),
519+
_point({**_EP_WEATHER}, 0.009),
520+
]
521+
522+
dense_resp = _make_qdrant_dense_response(dense_points)
523+
hybrid_resp = _make_qdrant_hybrid_response(hybrid_points)
524+
count_resp = _make_count_response(10)
525+
526+
client = AsyncMock()
527+
client.get = AsyncMock(return_value=count_resp)
528+
client.post = AsyncMock(side_effect=[dense_resp, hybrid_resp])
529+
530+
searcher = _make_searcher(client)
531+
532+
# asyncio.to_thread is called twice:
533+
# 1st call → _get_query_embedding → must return a valid embedding vector
534+
# 2nd call → disambiguator.forward → must return None ("none" response)
535+
precomputed = [0.1] * 10
536+
_call_count = 0
537+
538+
async def _to_thread_side_effect(fn: Any, *args: Any, **kwargs: Any) -> Any:
539+
nonlocal _call_count
540+
_call_count += 1
541+
if _call_count == 1:
542+
return precomputed # embedding call
543+
return None # disambiguator call → rejects all candidates
544+
545+
with patch(
546+
"tool_classifier.api_semantic_searcher.asyncio.to_thread",
547+
side_effect=_to_thread_side_effect,
548+
):
549+
results = await searcher.search("holidays AND weather")
550+
551+
# Must return exactly one result — the top cosine candidate
552+
assert len(results) == 1
553+
top = results[0]
554+
# Top candidate by cosine score is ep-holidays
555+
assert top.endpoint_id == "ep-holidays"
556+
# NOT llm_validated — disambiguator explicitly rejected it
557+
assert top.llm_validated is False
558+
# multi_intent_hint signals the classifier to try IntentDecomposer
559+
assert top.multi_intent_hint is True
560+
510561

511562
class TestSearchBelowThreshold:
512563
@pytest.mark.asyncio

0 commit comments

Comments
 (0)