新增功能

This commit is contained in:
Your Name
2026-07-28 09:46:53 +08:00
parent 45b3bc0852
commit 980795b4da
57 changed files with 18421 additions and 1051 deletions
+321
View File
@@ -0,0 +1,321 @@
from __future__ import annotations
import io
import json
import unittest
from types import SimpleNamespace
from unittest import mock
import grok_direct_chat
class _Manager:
def __init__(self, *, backend: str = "dify"):
self.backend = backend
def load_ai_settings(self):
return {
"GROK_MODEL_ENABLED": True,
"GROK_API_KEY": "app-test-key",
"GROK_CUSTOMER_SERVICE_TIMEOUT": 180,
"GROK_DIFY_INPUTS": {"tenant": "test"},
}
def model_profile(self, _settings):
return SimpleNamespace(
compatible=True,
reason="",
model="self-model",
base_url="https://model.example/v1",
api_backend=self.backend,
auth_scheme="bearer",
temperature=0.2,
max_completion_tokens=512,
)
class _StreamResponse(io.BytesIO):
def __init__(self, events: list[object], *, content_type: str = "text/event-stream"):
body_parts: list[bytes] = []
for event in events:
if isinstance(event, bytes):
body_parts.append(event)
elif event == "[DONE]":
body_parts.append(b"data: [DONE]\n\n")
else:
body_parts.append(
(
"data: "
+ json.dumps(event, ensure_ascii=False)
+ "\n\n"
).encode("utf-8")
)
super().__init__(b"".join(body_parts))
self.status = 200
self.headers = {"Content-Type": content_type}
class ChatRouteTests(unittest.TestCase):
def test_normal_language_defaults_to_direct(self):
for text in (
"你好",
"你是什么模型",
"解释一下量子纠缠",
"帮我润色这段文字",
"?",
):
with self.subTest(text=text):
self.assertEqual(grok_direct_chat.classify_chat_route(text), "direct")
def test_live_or_executable_work_uses_agent(self):
for text in (
"帮我查询郑州天气",
"查天气",
"查一下最新新闻",
"打开这个网页 https://example.com",
"运行项目里的测试脚本",
"修改这个文件",
"给企业微信联系人张三发送消息",
"使用 Agent 调用 MCP 工具",
"@agent 处理这个任务",
):
with self.subTest(text=text):
self.assertEqual(grok_direct_chat.classify_chat_route(text), "agent")
def test_discussion_about_tools_does_not_start_agent(self):
self.assertEqual(
grok_direct_chat.classify_chat_route("MCP 是什么意思?"),
"direct",
)
self.assertEqual(
grok_direct_chat.classify_chat_route("解释一下这段代码"),
"direct",
)
def test_short_follow_up_inherits_agent_route(self):
self.assertEqual(
grok_direct_chat.classify_chat_route("明天呢?", last_route="agent"),
"agent",
)
self.assertEqual(
grok_direct_chat.classify_chat_route("", last_route="direct"),
"direct",
)
def test_identity_reply_uses_actual_configured_model(self):
result = grok_direct_chat.direct_chat(
"你是哪个模型",
manager=_Manager(),
)
self.assertIn("self-model", result.text)
self.assertNotIn("xAI", result.text)
@mock.patch("grok_direct_chat._post_json")
def test_dify_direct_chat_reuses_conversation(self, post):
post.return_value = {
"answer": "直接回复",
"conversation_id": "dify-conv-1",
}
result = grok_direct_chat.direct_chat(
"继续说明",
history=[{"role": "assistant", "content": "上一轮"}],
conversation_id="dify-conv-old",
manager=_Manager(),
)
self.assertEqual(result.text, "直接回复")
self.assertEqual(result.conversation_id, "dify-conv-1")
args, kwargs = post.call_args
self.assertEqual(args[0], "https://model.example/v1/chat-messages")
self.assertEqual(kwargs["payload"]["conversation_id"], "dify-conv-old")
self.assertEqual(kwargs["payload"]["response_mode"], "blocking")
self.assertLessEqual(kwargs["timeout"], 90)
@mock.patch("grok_direct_chat._post_json")
def test_chat_completions_sends_bounded_history(self, post):
post.return_value = {
"choices": [{"message": {"content": "普通回复"}}]
}
result = grok_direct_chat.direct_chat(
"当前问题",
history=[{"role": "user", "content": "历史问题"}],
manager=_Manager(backend="chat_completions"),
)
self.assertEqual(result.text, "普通回复")
payload = post.call_args.kwargs["payload"]
self.assertEqual(payload["messages"][-1]["content"], "当前问题")
self.assertEqual(payload["messages"][-2]["content"], "历史问题")
self.assertFalse(payload["stream"])
class DirectChatStreamingTests(unittest.TestCase):
def _stream(self, backend: str, events: list[object]):
updates: list[tuple[str, bool]] = []
response = _StreamResponse(events)
with mock.patch.object(
grok_direct_chat._HTTP_OPENER,
"open",
return_value=response,
) as opened:
result = grok_direct_chat.stream_direct_chat(
"请回答",
on_update=lambda text, replace: updates.append((text, replace)),
manager=_Manager(backend=backend),
)
request = opened.call_args.args[0]
payload = json.loads(request.data.decode("utf-8"))
return result, updates, payload, request
def test_dify_streams_deltas_and_conversation_id(self):
result, updates, payload, request = self._stream(
"dify",
[
{"event": "message", "answer": "", "conversation_id": "c1"},
{"event": "message", "answer": "", "conversation_id": "c1"},
{"event": "message_end", "conversation_id": "c1"},
],
)
self.assertEqual(result.text, "你好")
self.assertEqual(result.conversation_id, "c1")
self.assertEqual(updates, [("", False), ("", False)])
self.assertEqual(payload["response_mode"], "streaming")
self.assertEqual(request.get_header("Accept"), "text/event-stream")
def test_dify_message_replace_replaces_full_text(self):
result, updates, _payload, _request = self._stream(
"dify",
[
{"event": "agent_message", "answer": "草稿"},
{"event": "message_replace", "answer": "最终答案"},
{"event": "message_end"},
],
)
self.assertEqual(result.text, "最终答案")
self.assertEqual(updates[-1], ("最终答案", True))
def test_chat_completions_streams_until_done(self):
result, updates, payload, _request = self._stream(
"chat_completions",
[
{"choices": [{"delta": {"role": "assistant"}, "finish_reason": None}]},
{"choices": [{"delta": {"content": "A"}, "finish_reason": None}]},
{"choices": [{"delta": {"content": "B"}, "finish_reason": "stop"}]},
"[DONE]",
],
)
self.assertEqual(result.text, "AB")
self.assertEqual(updates, [("A", False), ("B", False)])
self.assertTrue(payload["stream"])
def test_responses_streams_until_completed(self):
result, updates, payload, _request = self._stream(
"responses",
[
{"type": "response.output_text.delta", "delta": ""},
{"type": "response.output_text.delta", "delta": ""},
{"type": "response.completed", "response": {"status": "completed"}},
],
)
self.assertEqual(result.text, "甲乙")
self.assertEqual(updates, [("", False), ("", False)])
self.assertTrue(payload["stream"])
def test_anthropic_streams_text_until_message_stop(self):
result, updates, payload, request = self._stream(
"messages",
[
{"type": "message_start", "message": {"content": []}},
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "Hello"}},
{"type": "content_block_stop", "index": 0},
{"type": "message_delta", "delta": {"stop_reason": "end_turn"}},
{"type": "message_stop"},
],
)
self.assertEqual(result.text, "Hello")
self.assertEqual(updates, [("Hello", False)])
self.assertTrue(payload["stream"])
self.assertEqual(request.headers["Anthropic-version"], "2023-06-01")
def test_partial_eof_is_rejected(self):
response = _StreamResponse(
[{"type": "response.output_text.delta", "delta": "残缺"}]
)
with mock.patch.object(
grok_direct_chat._HTTP_OPENER,
"open",
return_value=response,
):
with self.assertRaisesRegex(
grok_direct_chat.DirectChatError,
"未正常结束",
):
grok_direct_chat.stream_direct_chat(
"请回答",
on_update=lambda _text, _replace: None,
manager=_Manager(backend="responses"),
)
def test_chat_completions_requires_finish_reason(self):
response = _StreamResponse(
[
{"choices": [{"delta": {"content": "文本"}, "finish_reason": None}]},
"[DONE]",
]
)
with mock.patch.object(
grok_direct_chat._HTTP_OPENER,
"open",
return_value=response,
):
with self.assertRaisesRegex(
grok_direct_chat.DirectChatError,
"finish_reason",
):
grok_direct_chat.stream_direct_chat(
"请回答",
on_update=lambda _text, _replace: None,
manager=_Manager(backend="chat_completions"),
)
def test_cancel_before_open(self):
cancellation = grok_direct_chat.DirectChatCancellation()
cancellation.cancel()
with mock.patch.object(grok_direct_chat._HTTP_OPENER, "open") as opened:
with self.assertRaises(grok_direct_chat.DirectChatCancelled):
grok_direct_chat.stream_direct_chat(
"请回答",
on_update=lambda _text, _replace: None,
manager=_Manager(backend="responses"),
cancellation=cancellation,
)
opened.assert_not_called()
def test_cancel_during_stream_closes_response(self):
cancellation = grok_direct_chat.DirectChatCancellation()
response = _StreamResponse(
[
{"type": "response.output_text.delta", "delta": "第一段"},
{"type": "response.completed", "response": {"status": "completed"}},
]
)
def cancel_after_first(_text, _replace):
cancellation.cancel()
with mock.patch.object(
grok_direct_chat._HTTP_OPENER,
"open",
return_value=response,
):
with self.assertRaises(grok_direct_chat.DirectChatCancelled):
grok_direct_chat.stream_direct_chat(
"请回答",
on_update=cancel_after_first,
manager=_Manager(backend="responses"),
cancellation=cancellation,
)
self.assertTrue(response.closed)
if __name__ == "__main__":
unittest.main()