ecb52e9978
A server is a row with a URL; its tools are discovered by a button and cached, then offered beside the built-in ones. Written by hand rather than taken from the reference SDK, because that SDK's transport does its own connecting -- and the one thing that must not be bypassed is check_url on every hop. Owning the transport is the point; the framing beside it is the small part. Sessions are per call: initialize, initialized, the call, a best-effort DELETE. Caching one wants an owner, a TTL, eviction, a lock and a shutdown hook, and the server may expire it under all of that anyway -- ToolContext is a session-free snapshot precisely so nothing in a tool holds live state. A server's names and descriptions reach the model as instructions and are bounded before they do; what it returns is escaped preformatted text, never markdown. Tools are namespaced per server, so two servers exposing "search" do not collide and neither shadows a built-in. Also: a round's calls now run together under a semaphore, results indexed so each tool turn stays paired with its call, and generation.status names what is running -- a remote tool is latency-bound, and a silent pause is what a hang looks like. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
333 lines
12 KiB
Python
333 lines
12 KiB
Python
"""The tool loop: one reply, several requests."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from lembas.db.models import ROLE_ASSISTANT, Chat, Connection, Message, Model
|
|
from lembas.services import generation as generation_service
|
|
from lembas.services import settings_store
|
|
from lembas.services import tools as tools_service
|
|
from lembas.services.search.base import SearchResult
|
|
|
|
|
|
def _chat_with_tools(db, user_id):
|
|
connection = Connection(name="c", base_url="http://127.0.0.1:1", api_key_encrypted="")
|
|
db.add(connection)
|
|
db.commit()
|
|
db.add(Model(connection_id=connection.id, model_id="m", capabilities_json={"tools": True}))
|
|
db.commit()
|
|
chat = Chat(user_id=user_id, model_id="m", connection_id=connection.id)
|
|
db.add(chat)
|
|
db.commit()
|
|
db.add(Message(chat_id=chat.id, role="user", content="What is a mallorn?", complete=True))
|
|
db.commit()
|
|
assistant = Message(chat_id=chat.id, role=ROLE_ASSISTANT, content="", complete=False)
|
|
db.add(assistant)
|
|
db.commit()
|
|
return chat.id, assistant.id
|
|
|
|
|
|
def _tool_call_chunk(name: str, arguments: str) -> dict:
|
|
return {
|
|
"choices": [
|
|
{"delta": {"tool_calls": [{"index": 0, "id": "c1", "function": {
|
|
"name": name, "arguments": arguments}}]}}
|
|
]
|
|
}
|
|
|
|
|
|
def _text_chunk(text: str) -> dict:
|
|
return {"choices": [{"delta": {"content": text}}]}
|
|
|
|
|
|
def _stub_stream(rounds, seen_payloads):
|
|
"""A stream_chat that returns a different scripted round each time."""
|
|
|
|
async def stream_chat(_endpoint, payload):
|
|
seen_payloads.append(payload)
|
|
for chunk in rounds[min(len(seen_payloads) - 1, len(rounds) - 1)]:
|
|
yield chunk
|
|
|
|
return stream_chat
|
|
|
|
|
|
async def test_a_tool_call_produces_a_second_request(db, user_id, monkeypatch):
|
|
"""The whole point: one reply, two round trips, with the search result in
|
|
the second one's messages."""
|
|
settings_store.update(db, {"enabled": True}, key=settings_store.SEARCH)
|
|
chat_id, message_id = _chat_with_tools(db, user_id)
|
|
|
|
async def fake_search(_config, _query, *, limit=None):
|
|
return [SearchResult("Mallorn", "https://tolkien.test/mallorn", "A golden tree.")]
|
|
|
|
monkeypatch.setattr("lembas.services.search.run", fake_search)
|
|
|
|
payloads = []
|
|
monkeypatch.setattr(
|
|
generation_service,
|
|
"stream_chat",
|
|
_stub_stream(
|
|
[
|
|
[_tool_call_chunk("web_search", '{"query": "mallorn"}')],
|
|
[_text_chunk("A mallorn is a golden tree.")],
|
|
],
|
|
payloads,
|
|
),
|
|
)
|
|
monkeypatch.setattr(
|
|
"lembas.services.chat.generate_title", _never_called_title
|
|
)
|
|
|
|
generation = generation_service.Generation(chat_id=chat_id, message_id=message_id)
|
|
await generation_service._run(generation)
|
|
|
|
assert len(payloads) == 2, "the model asked for a tool, so it must be asked again"
|
|
assert generation.text == "A mallorn is a golden tree."
|
|
|
|
# The second request carries the assistant's own call back, then the result.
|
|
followups = payloads[1]["messages"][-2:]
|
|
assert followups[0]["tool_calls"][0]["function"]["name"] == "web_search"
|
|
assert followups[1]["role"] == "tool"
|
|
assert "https://tolkien.test/mallorn" in followups[1]["content"]
|
|
|
|
# And the reader gets to see what it looked up.
|
|
assert generation.tool_events[0]["query"] == "mallorn"
|
|
assert generation.tool_events[0]["results"][0]["url"] == "https://tolkien.test/mallorn"
|
|
|
|
|
|
async def test_the_tools_array_is_absent_without_the_capability(db, user_id, monkeypatch):
|
|
chat_id, message_id = _chat_with_tools(db, user_id)
|
|
# Search enabled, but the model is not marked as supporting tools.
|
|
settings_store.update(db, {"enabled": True}, key=settings_store.SEARCH)
|
|
model = db.query(Model).first()
|
|
model.capabilities_json = {}
|
|
db.commit()
|
|
|
|
payloads = []
|
|
monkeypatch.setattr(
|
|
generation_service, "stream_chat", _stub_stream([[_text_chunk("hi")]], payloads)
|
|
)
|
|
monkeypatch.setattr("lembas.services.chat.generate_title", _never_called_title)
|
|
|
|
await generation_service._run(
|
|
generation_service.Generation(chat_id=chat_id, message_id=message_id)
|
|
)
|
|
assert "tools" not in payloads[0]
|
|
|
|
|
|
async def test_text_before_a_tool_call_is_kept(db, user_id, monkeypatch):
|
|
"""A model that narrates what it is about to look up must not lose that
|
|
when the results come back."""
|
|
settings_store.update(db, {"enabled": True}, key=settings_store.SEARCH)
|
|
chat_id, message_id = _chat_with_tools(db, user_id)
|
|
|
|
monkeypatch.setattr(
|
|
"lembas.services.search.run", _empty_search
|
|
)
|
|
monkeypatch.setattr(
|
|
generation_service,
|
|
"stream_chat",
|
|
_stub_stream(
|
|
[
|
|
[
|
|
_text_chunk("Let me look that up. "),
|
|
_tool_call_chunk("web_search", '{"query": "mallorn"}'),
|
|
],
|
|
[_text_chunk("Nothing found.")],
|
|
],
|
|
[],
|
|
),
|
|
)
|
|
monkeypatch.setattr("lembas.services.chat.generate_title", _never_called_title)
|
|
|
|
generation = generation_service.Generation(chat_id=chat_id, message_id=message_id)
|
|
await generation_service._run(generation)
|
|
assert generation.text == "Let me look that up. Nothing found."
|
|
|
|
|
|
async def test_a_model_that_only_ever_calls_tools_is_stopped(db, user_id, monkeypatch):
|
|
"""Otherwise a small model that has decided searching is the answer keeps
|
|
searching until the context runs out, at a full request each time."""
|
|
settings_store.update(db, {"enabled": True}, key=settings_store.SEARCH)
|
|
chat_id, message_id = _chat_with_tools(db, user_id)
|
|
|
|
monkeypatch.setattr("lembas.services.search.run", _empty_search)
|
|
|
|
payloads = []
|
|
monkeypatch.setattr(
|
|
generation_service,
|
|
"stream_chat",
|
|
_stub_stream([[_tool_call_chunk("web_search", '{"query": "x"}')]], payloads),
|
|
)
|
|
monkeypatch.setattr("lembas.services.chat.generate_title", _never_called_title)
|
|
|
|
generation = generation_service.Generation(chat_id=chat_id, message_id=message_id)
|
|
await generation_service._run(generation)
|
|
|
|
assert len(payloads) == tools_service.MAX_ROUNDS + 1
|
|
# Recorded rather than silently dropped: an answer that stops here has to
|
|
# be explicable.
|
|
assert generation.tool_events[-1]["status"] == "error"
|
|
|
|
|
|
async def test_tool_activity_is_stored_with_the_message(db, user_id, monkeypatch):
|
|
settings_store.update(db, {"enabled": True}, key=settings_store.SEARCH)
|
|
chat_id, message_id = _chat_with_tools(db, user_id)
|
|
|
|
async def fake_search(_config, _query, *, limit=None):
|
|
return [SearchResult("Mallorn", "https://tolkien.test/m", "A tree.")]
|
|
|
|
monkeypatch.setattr("lembas.services.search.run", fake_search)
|
|
monkeypatch.setattr(
|
|
generation_service,
|
|
"stream_chat",
|
|
_stub_stream(
|
|
[
|
|
[_tool_call_chunk("web_search", '{"query": "mallorn"}')],
|
|
[_text_chunk("Done.")],
|
|
],
|
|
[],
|
|
),
|
|
)
|
|
monkeypatch.setattr("lembas.services.chat.generate_title", _never_called_title)
|
|
|
|
await generation_service._run(
|
|
generation_service.Generation(chat_id=chat_id, message_id=message_id)
|
|
)
|
|
|
|
stored = db.get(Message, message_id)
|
|
db.refresh(stored)
|
|
assert stored.tool_calls_json[0]["query"] == "mallorn"
|
|
assert stored.complete is True
|
|
|
|
|
|
async def _empty_search(_config, _query, *, limit=None):
|
|
return []
|
|
|
|
|
|
async def _never_called_title(*_args, **_kwargs):
|
|
"""Auto-titling makes its own request; these tests are about the tool loop."""
|
|
return "A title"
|
|
|
|
|
|
# --- Progress and concurrency ------------------------------------------------
|
|
async def test_the_status_names_the_running_tool_and_is_cleared(db, user_id, monkeypatch):
|
|
"""A remote tool can take seconds with nothing streaming, and a silent
|
|
pause is exactly what a hang looks like."""
|
|
settings_store.update(db, {"enabled": True}, key=settings_store.SEARCH)
|
|
chat_id, message_id = _chat_with_tools(db, user_id)
|
|
|
|
seen: list[str] = []
|
|
|
|
async def fake_search(_config, _query, *, limit=None):
|
|
seen.append(generation.status)
|
|
return []
|
|
|
|
monkeypatch.setattr("lembas.services.search.run", fake_search)
|
|
monkeypatch.setattr(
|
|
generation_service,
|
|
"stream_chat",
|
|
_stub_stream(
|
|
[[_tool_call_chunk("web_search", '{"query": "mallorn"}')], [_text_chunk("Done.")]],
|
|
[],
|
|
),
|
|
)
|
|
|
|
generation = generation_service.Generation(chat_id=chat_id, message_id=message_id)
|
|
await generation_service._run(generation)
|
|
|
|
assert seen == ["Running web_search…"]
|
|
assert generation.status == "", "and it is cleared once they are done"
|
|
|
|
|
|
async def test_results_stay_paired_with_their_calls_when_run_together(db, user_id, monkeypatch):
|
|
"""Indexed rather than appended as they finish: an endpoint matching on
|
|
tool_call_id would otherwise pair the right id with the wrong content."""
|
|
import asyncio
|
|
|
|
settings_store.update(db, {"enabled": True}, key=settings_store.SEARCH)
|
|
chat_id, message_id = _chat_with_tools(db, user_id)
|
|
|
|
async def slow_first(_config, query, *, limit=None):
|
|
# The first call finishes last, which is the whole point of the test.
|
|
await asyncio.sleep(0.02 if query == "first" else 0)
|
|
return [SearchResult(f"result for {query}", f"https://t.test/{query}", "")]
|
|
|
|
monkeypatch.setattr("lembas.services.search.run", slow_first)
|
|
|
|
two_calls = {
|
|
"choices": [
|
|
{"delta": {"tool_calls": [
|
|
{"index": 0, "id": "a", "function": {
|
|
"name": "web_search", "arguments": '{"query": "first"}'}},
|
|
{"index": 1, "id": "b", "function": {
|
|
"name": "web_search", "arguments": '{"query": "second"}'}},
|
|
]}}
|
|
]
|
|
}
|
|
payloads: list[dict] = []
|
|
monkeypatch.setattr(
|
|
generation_service,
|
|
"stream_chat",
|
|
_stub_stream([[two_calls], [_text_chunk("Done.")]], payloads),
|
|
)
|
|
|
|
generation = generation_service.Generation(chat_id=chat_id, message_id=message_id)
|
|
await generation_service._run(generation)
|
|
|
|
turns = [m for m in payloads[1]["messages"] if m.get("role") == "tool"]
|
|
assert [turn["tool_call_id"] for turn in turns] == ["a", "b"]
|
|
assert "first" in turns[0]["content"] and "second" in turns[1]["content"]
|
|
# And the transcript keeps the same order.
|
|
assert [event["query"] for event in generation.tool_events] == ["first", "second"]
|
|
|
|
|
|
async def test_a_custom_tool_runs_inside_the_loop(db, user_id, monkeypatch, mock_http):
|
|
"""End to end: a row becomes an offered schema, the model calls it, and the
|
|
result comes back in the next request's messages."""
|
|
import httpx
|
|
|
|
from lembas.db.models import CustomTool
|
|
|
|
monkeypatch.setattr(
|
|
"socket.getaddrinfo", lambda *a, **k: [(2, 1, 6, "", ("93.184.216.34", 80))]
|
|
)
|
|
mock_http(lambda _r: httpx.Response(200, json={"summary": "Sunny in Minas Tirith."}))
|
|
|
|
db.add(
|
|
CustomTool(
|
|
slug="weather",
|
|
name="Weather",
|
|
description="Look up the weather.",
|
|
url_template="https://api.test/{{city}}",
|
|
parameters_json={"type": "object", "properties": {"city": {"type": "string"}}},
|
|
response_mode="json",
|
|
response_path="summary",
|
|
)
|
|
)
|
|
db.commit()
|
|
|
|
chat_id, message_id = _chat_with_tools(db, user_id)
|
|
payloads: list[dict] = []
|
|
monkeypatch.setattr(
|
|
generation_service,
|
|
"stream_chat",
|
|
_stub_stream(
|
|
[
|
|
[_tool_call_chunk("weather", '{"city": "Minas Tirith"}')],
|
|
[_text_chunk("It is sunny.")],
|
|
],
|
|
payloads,
|
|
),
|
|
)
|
|
|
|
generation = generation_service.Generation(chat_id=chat_id, message_id=message_id)
|
|
await generation_service._run(generation)
|
|
|
|
offered = {tool["function"]["name"] for tool in payloads[0]["tools"]}
|
|
assert "weather" in offered
|
|
|
|
tool_turns = [m for m in payloads[1]["messages"] if m.get("role") == "tool"]
|
|
assert tool_turns[0]["content"] == "Sunny in Minas Tirith."
|
|
assert generation.tool_events[0]["kind"] == "custom"
|
|
assert generation.tool_events[0]["label"] == "Weather"
|