# -*- 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()