diff --git a/docs/self-evolution.md b/docs/self-evolution.md index 9c41e3e..2ab94cb 100644 --- a/docs/self-evolution.md +++ b/docs/self-evolution.md @@ -9,7 +9,7 @@ Historical verification snapshot: `51be33361422e55e1f2f00c33a0e0f8c56132a91` (the post-#54 `main` revision, captured before this #55 documentation-only update). Snapshot date: 2026-09-04. -Current repository test count at this snapshot: **192 unittest cases**. +Current repository test count at this snapshot: **193 unittest cases**. ## Verified surface diff --git a/main.py b/main.py index 5ba68f4..b4f5467 100644 --- a/main.py +++ b/main.py @@ -3930,6 +3930,7 @@ class DeepSeekDSMLFilter: r"invoke|parameter|calls|tool_calls|function_calls)\s*>", re.IGNORECASE, ) + _GENERIC_KINDS = ("invoke", "parameter", "calls", "tool_calls", "function_calls") _ACTION_MARKER_RE = _ACTION_MARKER_RE _CONTROL_PREFIXES = tuple( item[:length].lower() @@ -3966,6 +3967,21 @@ def _close_for(cls, match: re.Match[str]) -> str: @classmethod def _control_partial_suffix_length(cls, value: str) -> int: + start = value.rfind("<") + if start >= 0: + suffix = value[start:] + if ">" not in suffix: + remainder = suffix[1:] + if remainder.startswith("/"): + remainder = remainder[1:] + remainder = remainder.lstrip() + if not remainder: + return len(suffix) + name_match = re.match(r"[A-Za-z_][A-Za-z0-9_]*", remainder) + if name_match: + name = name_match.group(0).lower() + if any(kind.startswith(name) for kind in cls._GENERIC_KINDS): + return len(suffix) lowered = value.lower() for length in range(min(len(value), max(map(len, cls._CONTROL_PREFIXES))), 0, -1): if lowered[-length:] in cls._CONTROL_PREFIXES: diff --git a/tests/test_streaming.py b/tests/test_streaming.py index dc48f22..26987db 100644 --- a/tests/test_streaming.py +++ b/tests/test_streaming.py @@ -107,6 +107,16 @@ def chat_stream(self, _messages, _model=None): ) +class _SplitGenericInvokeProvider: + config = ProviderConfig(provider="deepseek", model_simple="deepseek-v4-flash") + marker = '< invoke name="read_file">read_file("README.md")' + + def chat_stream(self, _messages, _model=None): + yield "Visible progress. " + yield from self.marker + yield " after marker." + + class StreamingEndpointTests(unittest.TestCase): def test_first_sse_content_arrives_before_provider_stream_completes(self): session_id = "stream-timing-regression" @@ -361,6 +371,62 @@ def text(item): ] self.assertEqual(assistant_messages[-1]["content"], expected) + def test_generic_invoke_tags_split_at_every_boundary_stay_out_of_sse(self): + marker = _SplitGenericInvokeProvider.marker + expected = "Visible progress. after marker." + for split in range(1, len(marker)): + stream_filter = main.DeepSeekDSMLFilter() + content = stream_filter.feed("Visible progress. " + marker[:split]) + content += stream_filter.feed(marker[split:] + " after marker.", final=True) + self.assertEqual(content, expected, split) + + session_id = "stream-generic-invoke-regression" + provider = _SplitGenericInvokeProvider() + + async def exercise(): + response = await server.api_chat_stream(_Request({ + "message": "read the README", "session_id": session_id, + })) + iterator = response.body_iterator.__aiter__() + frames = [] + try: + while True: + frames.append(await asyncio.wait_for(anext(iterator), timeout=2)) + except StopAsyncIteration: + return frames + + with tempfile.TemporaryDirectory(prefix="openkyrozen-generic-invoke-") as directory: + memory = MemoryBank(Path(directory) / "state.sqlite3", workspace_id="generic-invoke-test") + original_sessions = server._sessions + server._sessions = {} + try: + with (patch.object(server._agent, "memory_bank", memory), + patch.object(server._agent, "llm_provider", provider), + patch.object(server._agent, "DEEPSEEK_MODEL", "deepseek-v4-flash"), + patch.object(server._agent, "_chat_turn", side_effect=lambda message, **_: ( + main._call_llm_with_spinner([{"role": "user", "content": message}]) + ))): + frames = asyncio.run(exercise()) + finally: + server._sessions = original_sessions + + def text(item): + return item.decode() if isinstance(item, bytes) else item + + payloads = [json.loads(text(frame).split("data: ", 1)[1].splitlines()[0]) + for frame in frames + if text(frame).startswith("data: ") and text(frame) != "data: [DONE]\n\n"] + content = "".join(item["chunk"] for item in payloads if item.get("event") == "content") + self.assertEqual(content, expected) + self.assertNotRegex(content, r"<\s*/?\s*(?:invoke|parameter|calls|tool_calls|function_calls)\b") + self.assertNotIn("read_file(\"README.md\")", content) + assistant_messages = [ + event["payload"] for event in memory.store.list_events( + "session.message", workspace_id="generic-invoke-test", session_id=session_id, + ) if event["payload"].get("role") == "assistant" + ] + self.assertEqual(assistant_messages[-1]["content"], expected) + if __name__ == "__main__": unittest.main()