From 2f1ab85aec314c524e0831fe4cdb62ea5cd3f77b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ey=C3=BCp=20Can=20Akman?= Date: Tue, 21 Apr 2026 11:36:56 +0300 Subject: [PATCH] Fix Mistral empty tool_call_end flipping state machine to normal (#1151) --- mlx_lm/server.py | 5 +++-- tests/test_server.py | 27 +++++++++++++++++++++++++++ 2 files changed, 30 insertions(+), 2 deletions(-) diff --git a/mlx_lm/server.py b/mlx_lm/server.py index 685ae9a..29b7dd7 100644 --- a/mlx_lm/server.py +++ b/mlx_lm/server.py @@ -674,10 +674,11 @@ class ResponseGenerator: ts = tokenizer.tool_call_start_tokens te = tokenizer.tool_call_end_tokens transitions["normal"].append((ts, "tool")) - transitions["tool"] = [(te, "normal")] + transitions["tool"] = [(te, "normal")] if te else [] transitions["tool"].extend(common_stops) sequences[ts] = tokenizer.tool_call_start - sequences[te] = tokenizer.tool_call_end + if te: + sequences[te] = tokenizer.tool_call_end sm = SequenceStateMachine(transitions, initial=initial_state) if len(self._state_machine_cache) > 100: diff --git a/tests/test_server.py b/tests/test_server.py index 7139ca7..0ed8e84 100644 --- a/tests/test_server.py +++ b/tests/test_server.py @@ -277,6 +277,33 @@ class TestServer(unittest.TestCase): self.assertIn("id", response_body) self.assertIn("choices", response_body) + def test_make_state_machine_empty_tool_call_end(self): + class FakeTokenizer: + has_thinking = False + has_tool_calling = True + tool_call_start = "[TOOL_CALLS]" + tool_call_end = "" + tool_call_start_tokens = (100,) + tool_call_end_tokens = () + eos_token_ids = [2] + + def convert_ids_to_tokens(self, t): + return f"" + + sm, _ = self.response_generator._make_state_machine( + ("fake-empty-end", None, None), + FakeTokenizer(), + stop_words=[], + ) + state = sm.make_state() + state, _, s = sm.match(state, 100) + self.assertEqual(s, "tool") + for tok in [42, 43, 44]: + state, _, s = sm.match(state, tok) + self.assertEqual(s, "tool") + state, _, s = sm.match(state, 2) + self.assertIsNone(s) + def test_handle_models(self): url = f"http://localhost:{self.port}/v1/models" response = requests.get(url)