From 3cecefccdcf2220842aadd92961d78dc97e696ba Mon Sep 17 00:00:00 2001 From: onion-111 Date: Sat, 26 Sep 2026 23:41:33 +0800 Subject: [PATCH] [Feat] Support local decision models in MCP decide --- CHANGELOG.md | 2 + docs/skills.md | 9 +++- s1a/mcp_server.py | 27 ++++++++---- tests/test_mcp_server.py | 94 ++++++++++++++++++++++++++++++++++++++-- 4 files changed, 118 insertions(+), 14 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 21b9dd7..533ff21 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -6,6 +6,8 @@ The format follows [Keep a Changelog](https://keepachangelog.com/en/1.1.0/); ver ### Added +- The MCP `decide` tool accepts `model="jev"|"laya"|"cua"`, defaulting to `jev`. Local backends use their + optional extras and need no Jev API key. - `docs/benchmarks.md`: the Google Flights driver comparison rerun on 2026-09-23 from Poland, every arm three times on both decision backends, next to the baseline rows in one table; the 24 S1A records, as one archive, and the chart under `docs/results/flights/rerun-2026-09-23/`. diff --git a/docs/skills.md b/docs/skills.md index cbd7f1b..b1700a6 100644 --- a/docs/skills.md +++ b/docs/skills.md @@ -74,8 +74,13 @@ codex mcp add s1a -- uv run --project /path/to/system1-agents s1a-mcp `s1a-mcp` serves the same agents over stdio as three tools. `list_agents()` returns every agent with its front, its description and the flags `run_agent` accepts for it; an agent whose optional dependency is missing is listed as unavailable with the error. `run_agent(name, flags)` runs one agent with the flags of `s1a run ` and returns -its JSON object. `decide(state, options, rules)` answers one choice question with `jev`: the chosen key, a -probability per option, a confidence and the latency in ms. +its JSON object. `decide(state, options, rules, model="jev")` answers one choice question: the chosen key, a +probability per option, a confidence and the latency in ms. `model` accepts `jev`, `laya` or `cua`; callers that +omit it keep using Jev. For local decisions, install the matching extra in the server's checkout (`uv sync +--extra laya` or `uv sync --extra cua`) and pass `model="laya"` or `model="cua"`; no Jev API key is needed. +The first local call may download the checkpoint. Each call loads and closes its model; `ms` measures the +decision, not model loading. Agent runs and decisions are serialized, and model output stays off the stdio +protocol stream. ## Build a System 1 agent diff --git a/s1a/mcp_server.py b/s1a/mcp_server.py index d6f60dd..aa35f03 100644 --- a/s1a/mcp_server.py +++ b/s1a/mcp_server.py @@ -9,7 +9,7 @@ import logging import sys from collections.abc import AsyncIterator -from typing import Any +from typing import Any, Literal import s1a.console # first: routes the harness logs to files before openjiuwen loads and logs to the console from mcp.server.fastmcp import FastMCP @@ -25,7 +25,8 @@ "constraint puzzles or free-text generation. list_agents gives every agent's flags: a browser agent takes " "--model jev --goal '...' and needs a chat-model key (OPENAI_API_KEY or LLM_API_KEY plus MODEL_NAME), a Jev key " "(TYPESAFE_API_KEY or OPENROUTER_API_KEY) and Node for @playwright/mcp; a tool agent takes --model, --rethink and " - "--episodes; a rail takes --labelled-set. Runs go one at a time per server." + "--episodes; a rail takes --labelled-set. decide accepts model jev (default), laya or cua; local models need " + "their optional extra and no Jev key. Runs and decisions go one at a time per server." ) @@ -89,13 +90,21 @@ async def run_agent(name: str, flags: list[str]) -> dict[str, Any]: @server.tool() -async def decide(state: dict[str, Any], options: dict[str, str], rules: str) -> dict[str, Any]: - """One choice question over options you enumerate: the chosen key, a probability per option, a confidence, the latency in ms.""" - decision_model = build_model("jev") - try: - return await probe.pick(decision_model, state=state, options=options, rules=rules) - finally: - await decision_model.close() +async def decide( + state: dict[str, Any], options: dict[str, str], rules: str, model: Literal["jev", "laya", "cua"] = "jev" +) -> dict[str, Any]: + """One choice question: the chosen key, probabilities, confidence and decision latency in ms. + + Use jev over HTTP (default), or laya/cua locally after installing the matching extra. Local model loading + is excluded from the reported latency. + """ + async with _ONE_RUN: + with contextlib.redirect_stdout(sys.stderr): # local SDK loading must not write to the stdio protocol + decision_model = await asyncio.to_thread(build_model, model) + try: + return await probe.pick(decision_model, state=state, options=options, rules=rules) + finally: + await decision_model.close() def main() -> None: diff --git a/tests/test_mcp_server.py b/tests/test_mcp_server.py index 6d99215..737db84 100644 --- a/tests/test_mcp_server.py +++ b/tests/test_mcp_server.py @@ -16,7 +16,7 @@ from pathlib import Path from typing import Any from unittest import IsolatedAsyncioTestCase, TestCase, skipUnless -from unittest.mock import patch +from unittest.mock import AsyncMock, patch from mcp import ClientSession, StdioServerParameters from mcp.client.stdio import stdio_client @@ -24,7 +24,7 @@ from s1a import mcp_server from s1a import run as agents -from s1a.decision_models import JevModel, ScriptedTransport +from s1a.decision_models import JevModel, ScriptedModel, ScriptedTransport from s1a.tool import loop, series from support import COUNTER @@ -136,7 +136,7 @@ async def test_exposes_the_three_tools(self) -> None: async def test_decide_returns_the_validated_choice(self) -> None: transport = ScriptedTransport(choose="inc") - with patch.object(mcp_server, "build_model", lambda model_name, **kwargs: JevModel(transport)): + with patch.object(mcp_server, "build_model", return_value=JevModel(transport)) as build: async with create_connected_server_and_client_session(mcp_server.server) as session: result = await session.call_tool( "decide", {"state": {"n": 1}, "options": {"inc": "add one", "noop": "do nothing"}, "rules": "count"} @@ -145,6 +145,94 @@ async def test_decide_returns_the_validated_choice(self) -> None: answer = _rows(result) self.assertEqual((answer["choice"], answer["probabilities"]["inc"], answer["ms"]), ("inc", 1.0, 9)) self.assertEqual(transport.bodies[0]["questions"]["pick"]["instructions"]["rules"], "count") + build.assert_called_once_with("jev") + + async def test_decide_advertises_optional_model_choices(self) -> None: + async with create_connected_server_and_client_session(mcp_server.server) as session: + tools = await session.list_tools() + schema = next(tool.inputSchema for tool in tools.tools if tool.name == "decide") + self.assertEqual(schema["properties"]["model"]["enum"], ["jev", "laya", "cua"]) + self.assertEqual(schema["properties"]["model"]["default"], "jev") + self.assertNotIn("model", schema["required"]) + + async def test_decide_selects_and_closes_each_local_model_without_api_keys(self) -> None: + for model_name in ("laya", "cua"): + with self.subTest(model=model_name): + decision_model = ScriptedModel(choose="inc") + with ( + patch.dict(os.environ, NO_KEYS), + patch.object(mcp_server, "build_model", return_value=decision_model) as build, + patch.object(decision_model, "close", new_callable=AsyncMock) as close, + ): + async with create_connected_server_and_client_session(mcp_server.server) as session: + result = await session.call_tool( + "decide", + {"state": {"n": 1}, "options": {"inc": "add one"}, "rules": "count", "model": model_name}, + ) + self.assertFalse(result.isError) + self.assertEqual( + _rows(result), {"choice": "inc", "probabilities": {"inc": 1.0}, "confidence": 1.0, "ms": 9} + ) + build.assert_called_once_with(model_name) + close.assert_awaited_once() + + async def test_decide_rejects_unsupported_models_before_building_one(self) -> None: + with patch.object(mcp_server, "build_model") as build: + async with create_connected_server_and_client_session(mcp_server.server) as session: + for model_name in ("llm", "random", "rule", "unknown"): + with self.subTest(model=model_name): + result = await session.call_tool( + "decide", {"state": {}, "options": {"a": "a"}, "rules": "pick", "model": model_name} + ) + self.assertTrue(result.isError) + build.assert_not_called() + + async def test_decide_closes_a_local_model_when_the_decision_fails(self) -> None: + decision_model = ScriptedModel(error=RuntimeError("local inference failed")) + with ( + patch.object(mcp_server, "build_model", return_value=decision_model), + patch.object(decision_model, "close", new_callable=AsyncMock) as close, + ): + async with create_connected_server_and_client_session(mcp_server.server) as session: + result = await session.call_tool( + "decide", {"state": {}, "options": {"a": "a"}, "rules": "pick", "model": "cua"} + ) + self.assertTrue(result.isError) + self.assertIn("local inference failed", result.content[0].text) + close.assert_awaited_once() + + async def test_decide_reports_how_to_install_a_missing_local_model(self) -> None: + for model_name, module_name in (("laya", "laya"), ("cua", "cua_s1.nano")): + with ( + self.subTest(model=model_name), + patch.dict(os.environ, NO_KEYS), + patch.dict(sys.modules, {module_name: None}), + ): + async with create_connected_server_and_client_session(mcp_server.server) as session: + result = await session.call_tool( + "decide", {"state": {}, "options": {"a": "a"}, "rules": "pick", "model": model_name} + ) + self.assertTrue(result.isError) + self.assertIn(f"uv sync --extra {model_name}", result.content[0].text) + + async def test_local_model_loading_keeps_stdout_clean(self) -> None: + def build(model_name: str) -> ScriptedModel: + print("loading local weights") + return ScriptedModel(choose="a") + + stdout, stderr = io.StringIO(), io.StringIO() + with ( + patch.object(mcp_server, "build_model", build), + contextlib.redirect_stdout(stdout), + contextlib.redirect_stderr(stderr), + ): + async with create_connected_server_and_client_session(mcp_server.server) as session: + result = await session.call_tool( + "decide", {"state": {}, "options": {"a": "a"}, "rules": "pick", "model": "laya"} + ) + self.assertFalse(result.isError) + self.assertNotIn("loading local weights", stdout.getvalue()) + self.assertIn("loading local weights", stderr.getvalue()) async def test_list_agents_names_every_front_with_its_flags(self) -> None: async with create_connected_server_and_client_session(mcp_server.server) as session: