Files
kefu/wechat_rpa/test_dify_grok_adapter.py
2026-07-28 09:46:53 +08:00

791 lines
27 KiB
Python

# -*- coding: utf-8 -*-
"""Tests for the loopback Dify model adapter used by Grok Build."""
from __future__ import annotations
import base64
import json
import tempfile
import threading
import unittest
import urllib.error
import urllib.request
import uuid
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from dify_grok_adapter import (
ensure_dify_adapter,
stop_dify_adapter,
)
from grok_build_bridge import MODEL_API_KEY_ENV, GrokBuildManager
class _FakeDifyServer(ThreadingHTTPServer):
daemon_threads = True
def __init__(self):
super().__init__(("127.0.0.1", 0), _FakeDifyHandler)
self.requests: list[dict] = []
self.uploads: list[bytes] = []
self.responder = self._default_responder
self.emit_workflow_started = False
self.emit_workflow_finished = False
self.emit_malformed_sse_data = False
@staticmethod
def _envelope(prompt: str) -> dict:
start = "BEGIN_GROK_PROTOCOL_JSON\n"
end = "\nEND_GROK_PROTOCOL_JSON"
return json.loads(prompt.split(start, 1)[1].split(end, 1)[0])
def _default_responder(self, payload: dict) -> str:
envelope = self._envelope(str(payload.get("query") or ""))
tool_choice = envelope.get("tool_choice") or {}
if tool_choice.get("mode") == "function":
name = str(tool_choice.get("name") or "")
selected = next(
tool for tool in envelope["tools"] if tool["name"] == name
)
properties = selected["parameters"].get("properties") or {}
arguments = {
key: value["const"]
for key, value in properties.items()
if isinstance(value, dict) and "const" in value
}
return json.dumps(
{
"kind": "tool_calls",
"tool_calls": [
{
"name": name,
"arguments": arguments,
}
],
},
ensure_ascii=False,
)
if any(
message.get("role") == "tool"
for message in envelope.get("messages", [])
):
return json.dumps(
{"kind": "assistant", "content": "工具结果已收到"},
ensure_ascii=False,
)
return json.dumps(
{"kind": "assistant", "content": "你好"},
ensure_ascii=False,
)
class _FakeDifyHandler(BaseHTTPRequestHandler):
server: _FakeDifyServer
def log_message(self, _format: str, *_args: object) -> None:
return
def do_POST(self) -> None: # noqa: N802 - stdlib handler API
if self.path == "/v1/files/upload":
if self.headers.get("Authorization") != "Bearer upstream-secret":
self.send_error(401)
return
length = int(self.headers.get("Content-Length") or "0")
self.server.uploads.append(self.rfile.read(length))
body = json.dumps(
{
"id": f"upload-{len(self.server.uploads)}",
"name": "image.png",
}
).encode()
self.send_response(201)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
return
if self.path != "/v1/chat-messages":
self.send_error(404)
return
if self.headers.get("Authorization") != "Bearer upstream-secret":
body = json.dumps({"message": "unauthorized"}).encode()
self.send_response(401)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
return
length = int(self.headers.get("Content-Length") or "0")
payload = json.loads(self.rfile.read(length).decode("utf-8"))
self.server.requests.append(payload)
answer = self.server.responder(payload)
events = []
if self.server.emit_workflow_started:
events.append(
{
"event": "workflow_started",
"data": {"status": "running"},
}
)
events.extend([
{"event": "message", "answer": answer},
{
"event": "message_end",
"metadata": {"usage": {"total_tokens": 1}},
},
])
if self.server.emit_workflow_finished:
events.append(
{
"event": "workflow_finished",
"data": {"status": "succeeded"},
}
)
body = (
("data: {malformed-json\n\n" if self.server.emit_malformed_sse_data else "")
+ "".join(
"data: " + json.dumps(event, ensure_ascii=False) + "\n\n"
for event in events
)
).encode("utf-8")
self.send_response(200)
self.send_header("Content-Type", "text/event-stream")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
class DifyGrokAdapterTests(unittest.TestCase):
def setUp(self) -> None:
self.fake = _FakeDifyServer()
self.fake_thread = threading.Thread(
target=self.fake.serve_forever,
daemon=True,
)
self.fake_thread.start()
self.runtime_id = f"test-{uuid.uuid4()}"
self.info = ensure_dify_adapter(
self.runtime_id,
upstream_base_url=(
f"http://127.0.0.1:{self.fake.server_address[1]}"
"/v1/chat-messages"
),
api_key="upstream-secret",
model="private-model",
timeout=10,
inputs={"tenant": "医院"},
)
self.addCleanup(stop_dify_adapter, self.runtime_id)
self.addCleanup(self._stop_fake)
def _stop_fake(self) -> None:
self.fake.shutdown()
self.fake.server_close()
self.fake_thread.join(timeout=2)
def _request(
self,
payload: dict,
*,
token: str | None = None,
) -> tuple[int, str, str]:
request = urllib.request.Request(
f"{self.info.base_url}/chat/completions",
data=json.dumps(payload, ensure_ascii=False).encode("utf-8"),
headers={
"Content-Type": "application/json",
"Authorization": (
"Bearer "
+ (self.info.local_api_key if token is None else token)
),
},
method="POST",
)
try:
with urllib.request.urlopen(request, timeout=10) as response:
return (
int(response.status),
response.read().decode("utf-8"),
str(response.headers.get("Content-Type") or ""),
)
except urllib.error.HTTPError as exc:
try:
return (
int(exc.code),
exc.read().decode("utf-8"),
str(exc.headers.get("Content-Type") or ""),
)
finally:
exc.close()
@staticmethod
def _sse_values(body: str) -> list[object]:
values: list[object] = []
for line in body.splitlines():
if not line.startswith("data:"):
continue
raw = line[5:].strip()
values.append(raw if raw == "[DONE]" else json.loads(raw))
return values
def test_uses_random_local_token_and_normalizes_full_dify_endpoint(self) -> None:
self.assertNotEqual("upstream-secret", self.info.local_api_key)
self.assertEqual(
f"http://127.0.0.1:{self.fake.server_address[1]}/v1",
self.info.upstream_base_url,
)
status, _body, _content_type = self._request(
{
"model": "private-model",
"messages": [{"role": "user", "content": "你好"}],
"stream": True,
},
token="upstream-secret",
)
self.assertEqual(401, status)
def test_configuration_change_rotates_local_token(self) -> None:
unchanged = ensure_dify_adapter(
self.runtime_id,
upstream_base_url=self.info.upstream_base_url,
api_key="upstream-secret",
model="private-model",
timeout=10,
inputs={"tenant": "医院"},
)
self.assertEqual(self.info.local_api_key, unchanged.local_api_key)
changed = ensure_dify_adapter(
self.runtime_id,
upstream_base_url=self.info.upstream_base_url,
api_key="new-upstream-secret",
model="private-model",
timeout=10,
inputs={"tenant": "医院"},
)
self.assertNotEqual(self.info.port, changed.port)
self.assertNotEqual(self.info.local_api_key, changed.local_api_key)
status, _body, _content_type = self._request(
{
"model": "private-model",
"messages": [{"role": "user", "content": "你好"}],
"stream": True,
},
token=self.info.local_api_key,
)
self.assertEqual(200, status)
def test_streams_standard_text_chat_completion(self) -> None:
status, body, content_type = self._request(
{
"model": "private-model",
"messages": [{"role": "user", "content": "你好"}],
"stream": True,
"stream_options": {"include_usage": True},
}
)
self.assertEqual(200, status)
self.assertIn("text/event-stream", content_type)
values = self._sse_values(body)
text = "".join(
str(choice["delta"].get("content") or "")
for value in values
if isinstance(value, dict)
for choice in value.get("choices", [])
)
finishes = [
choice.get("finish_reason")
for value in values
if isinstance(value, dict)
for choice in value.get("choices", [])
if choice.get("finish_reason")
]
self.assertEqual("你好", text)
self.assertIn("stop", finishes)
usage_chunk = next(
value
for value in values
if isinstance(value, dict) and value.get("usage")
)
self.assertEqual(1, usage_chunk["usage"]["total_tokens"])
self.assertGreaterEqual(
usage_chunk["usage"]["completion_tokens"],
1,
)
self.assertEqual("[DONE]", values[-1])
self.assertEqual(
{"tenant": "医院"},
self.fake.requests[-1]["inputs"],
)
def test_uploads_data_uri_images_as_dify_files(self) -> None:
image_data = base64.b64encode(b"\x89PNG\r\n\x1a\nfake").decode()
status, _body, _content_type = self._request(
{
"model": "private-model",
"messages": [
{
"role": "tool",
"tool_call_id": "call_image",
"content": [
{
"type": "image_url",
"image_url": {
"url": f"data:image/png;base64,{image_data}"
},
}
],
}
],
"stream": True,
}
)
self.assertEqual(200, status)
self.assertEqual(1, len(self.fake.uploads))
self.assertEqual(
[
{
"type": "image",
"transfer_method": "local_file",
"upload_file_id": "upload-1",
}
],
self.fake.requests[-1]["files"],
)
def test_tool_call_and_followup_tool_result_round_trip(self) -> None:
schema = {
"type": "object",
"properties": {"query": {"type": "string", "const": "病历"}},
"required": ["query"],
"additionalProperties": False,
}
def tool_responder(payload: dict) -> str:
envelope = self.fake._envelope(payload["query"])
if any(
message.get("role") == "tool"
for message in envelope["messages"]
):
return json.dumps(
{"kind": "assistant", "content": "查询完成"},
ensure_ascii=False,
)
return json.dumps(
{
"kind": "tool_calls",
"tool_calls": [
{
"name": "search_records",
"arguments": {"query": "病历"},
}
],
},
ensure_ascii=False,
)
self.fake.responder = tool_responder
request_payload = {
"model": "private-model",
"messages": [{"role": "user", "content": "查病历"}],
"tools": [
{
"type": "function",
"function": {
"name": "search_records",
"description": "查询病历",
"parameters": schema,
},
}
],
"tool_choice": "auto",
"stream": True,
}
status, body, _content_type = self._request(request_payload)
self.assertEqual(200, status)
values = self._sse_values(body)
tool_delta = next(
choice["delta"]["tool_calls"][0]
for value in values
if isinstance(value, dict)
for choice in value.get("choices", [])
if choice.get("delta", {}).get("tool_calls")
)
self.assertEqual("search_records", tool_delta["function"]["name"])
call_id = tool_delta["id"]
request_payload["messages"] = [
{"role": "user", "content": "查病历"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": call_id,
"type": "function",
"function": {
"name": "search_records",
"arguments": '{"query":"病历"}',
},
}
],
},
{
"role": "tool",
"tool_call_id": call_id,
"content": "病历数据",
},
]
status, body, _content_type = self._request(request_payload)
self.assertEqual(200, status)
values = self._sse_values(body)
text = "".join(
str(choice["delta"].get("content") or "")
for value in values
if isinstance(value, dict)
for choice in value.get("choices", [])
)
self.assertEqual("查询完成", text)
def test_rejects_non_protocol_dify_text(self) -> None:
self.fake.responder = lambda _payload: "普通客服文本"
status, body, _content_type = self._request(
{
"model": "private-model",
"messages": [{"role": "user", "content": "你好"}],
"stream": True,
}
)
self.assertEqual(502, status)
self.assertIn("dify_protocol_error", body)
def test_rejects_unknown_tool_and_invalid_arguments(self) -> None:
self.fake.responder = lambda _payload: json.dumps(
{
"kind": "tool_calls",
"tool_calls": [
{
"name": "unknown",
"arguments": {},
}
],
}
)
status, body, _content_type = self._request(
{
"model": "private-model",
"messages": [{"role": "user", "content": "调用工具"}],
"tools": [
{
"type": "function",
"function": {
"name": "allowed",
"parameters": {
"type": "object",
"properties": {},
},
},
}
],
"stream": True,
}
)
self.assertEqual(502, status)
self.assertIn("dify_protocol_error", body)
def test_rejects_nonstandard_json_in_string_arguments(self) -> None:
self.fake.responder = lambda _payload: json.dumps(
{
"kind": "tool_calls",
"tool_calls": [
{
"name": "allowed",
"arguments": '{"value":NaN}',
}
],
}
)
status, body, _content_type = self._request(
{
"model": "private-model",
"messages": [{"role": "user", "content": "调用工具"}],
"tools": [
{
"type": "function",
"function": {
"name": "allowed",
"parameters": {
"type": "object",
"properties": {},
},
},
}
],
"stream": True,
}
)
self.assertEqual(502, status)
self.assertIn("dify_protocol_error", body)
def test_supports_local_schema_refs_and_rejects_remote_refs(self) -> None:
self.fake.responder = lambda _payload: json.dumps(
{
"kind": "tool_calls",
"tool_calls": [
{
"name": "search_records",
"arguments": {"query": "病历"},
}
],
},
ensure_ascii=False,
)
payload = {
"model": "private-model",
"messages": [{"role": "user", "content": "查病历"}],
"tools": [
{
"type": "function",
"function": {
"name": "search_records",
"parameters": {
"type": "object",
"$defs": {
"query": {
"type": "string",
"const": "病历",
}
},
"properties": {
"query": {"$ref": "#/$defs/query"}
},
"required": ["query"],
},
},
}
],
"stream": True,
}
status, _body, _content_type = self._request(payload)
self.assertEqual(200, status)
payload["tools"][0]["function"]["parameters"]["properties"]["query"] = {
"$ref": "https://example.test/schema.json"
}
status, body, _content_type = self._request(payload)
self.assertEqual(502, status)
self.assertIn("dify_protocol_error", body)
def test_chatflow_requires_workflow_finished(self) -> None:
self.fake.emit_workflow_started = True
status, body, _content_type = self._request(
{
"model": "private-model",
"messages": [{"role": "user", "content": "你好"}],
"stream": True,
}
)
self.assertEqual(502, status)
self.assertIn("dify_protocol_error", body)
self.fake.emit_workflow_finished = True
status, _body, _content_type = self._request(
{
"model": "private-model",
"messages": [{"role": "user", "content": "你好"}],
"stream": True,
}
)
self.assertEqual(200, status)
def test_rejects_malformed_dify_sse_data_instead_of_hiding_it(self) -> None:
self.fake.emit_malformed_sse_data = True
status, body, _content_type = self._request(
{
"model": "private-model",
"messages": [{"role": "user", "content": "你好"}],
"stream": True,
}
)
self.assertEqual(502, status)
self.assertIn("dify_protocol_error", body)
def test_numeric_bounds_disambiguate_one_of(self) -> None:
selected_value = {"value": 7}
self.fake.responder = lambda _payload: json.dumps(
{
"kind": "tool_calls",
"tool_calls": [
{
"name": "bounded",
"arguments": dict(selected_value),
}
],
}
)
payload = {
"model": "private-model",
"messages": [{"role": "user", "content": "选择数值"}],
"tools": [
{
"type": "function",
"function": {
"name": "bounded",
"parameters": {
"type": "object",
"properties": {
"value": {
"oneOf": [
{
"type": "integer",
"maximum": 5,
},
{
"type": "integer",
"minimum": 6,
"maximum": 10,
},
]
}
},
"required": ["value"],
},
},
}
],
"stream": True,
}
status, _body, _content_type = self._request(payload)
self.assertEqual(200, status)
selected_value["value"] = 11
status, body, _content_type = self._request(payload)
self.assertEqual(502, status)
self.assertIn("dify_protocol_error", body)
def test_bridge_writes_only_loopback_endpoint_and_local_token(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
settings_file = root / "ai_settings.json"
integration_file = root / "grok_build_settings.json"
settings = {
"GROK_MODEL_ENABLED": True,
"GROK_API_BASE": (
f"http://127.0.0.1:{self.fake.server_address[1]}/v1"
),
"GROK_API_KEY": "upstream-secret",
"GROK_MODEL": "private-model",
"GROK_API_BACKEND": "dify",
"GROK_AUTH_SCHEME": "auto",
"GROK_CONTEXT_WINDOW": 65536,
"GROK_MAX_TOKENS": 4096,
"GROK_TEMPERATURE": 0.2,
"GROK_CUSTOMER_SERVICE_TIMEOUT": 30,
"AI_MCP_SERVERS": [],
}
settings_file.write_text(
json.dumps(settings),
encoding="utf-8",
)
integration_file.write_text(
json.dumps(
{
"sync_mcp_servers": False,
"customer_service_tools": False,
}
),
encoding="utf-8",
)
manager = GrokBuildManager(
project_dir=root,
runtime_home=root / "runtime",
ai_settings_file=settings_file,
integration_settings_file=integration_file,
)
second_manager = GrokBuildManager(
project_dir=root,
runtime_home=root / "runtime",
ai_settings_file=settings_file,
integration_settings_file=integration_file,
)
self.addCleanup(
stop_dify_adapter,
str(manager.runtime_home),
)
concurrent_results: list[object] = []
concurrent_errors: list[Exception] = []
start = threading.Barrier(2)
def concurrent_sync(target: GrokBuildManager) -> None:
try:
start.wait(timeout=5)
concurrent_results.append(
target.sync_model_configuration()
)
except Exception as exc:
concurrent_errors.append(exc)
workers = [
threading.Thread(target=concurrent_sync, args=(target,))
for target in (manager, second_manager)
]
for worker in workers:
worker.start()
for worker in workers:
worker.join(timeout=10)
profile = manager.agent_model_profile()
sync = manager.sync_model_configuration()
environment = manager.runtime_environment(
include_model_key=True,
include_mcp_secrets=False,
)
probe = manager.probe_agent_model(force=True, timeout=10)
rendered = manager.user_config_file.read_text(encoding="utf-8")
live_status = manager.status()
stop_dify_adapter(str(manager.runtime_home))
stopped_status = manager.status()
self.assertTrue(profile.compatible)
self.assertFalse(concurrent_errors)
self.assertEqual(2, len(concurrent_results))
self.assertTrue(
all(
getattr(result, "compatible", False)
for result in concurrent_results
)
)
self.assertEqual("dify", profile.source_backend)
self.assertEqual("chat_completions", profile.api_backend)
self.assertTrue(profile.base_url.startswith("http://127.0.0.1:"))
self.assertEqual("dify", sync.api_backend)
self.assertEqual(
profile.source_base_url,
sync.base_url,
)
self.assertNotEqual(
"upstream-secret",
environment[MODEL_API_KEY_ENV],
)
self.assertNotIn("upstream-secret", rendered)
self.assertIn(profile.base_url, rendered)
self.assertTrue(probe.ok, probe.message)
self.assertEqual("dify", probe.api_backend)
self.assertTrue(live_status.adapter_live)
self.assertFalse(stopped_status.adapter_live)
if __name__ == "__main__":
unittest.main()