From 49afdc969229a5a07a09aba3806556d7a3b63dc1 Mon Sep 17 00:00:00 2001 From: "mark.tian" Date: Tue, 2 Dec 2025 18:31:51 +0800 Subject: [PATCH] =?UTF-8?q?=E5=A2=9E=E5=8A=A0remote=5Faddress=E5=92=8Cserv?= =?UTF-8?q?er=5Fip?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- app/database.py | 2 ++ app/routes.py | 11 +++++++++++ test_callback.py | 44 ++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 57 insertions(+) diff --git a/app/database.py b/app/database.py index c958d2f..874bbba 100644 --- a/app/database.py +++ b/app/database.py @@ -14,6 +14,8 @@ class CallbackLog(Base): id = Column(Integer, primary_key=True, autoincrement=True) site_id = Column(String(100), nullable=False, comment="站点ID") # 新增siteId字段 + remote_address = Column(String(45), nullable=True, comment="客户端IP地址") + server_ip = Column(String(45), nullable=True, comment="服务器IP地址") request_url = Column(String(500), nullable=False) request_headers = Column(JSON, nullable=False) request_body = Column(JSON, nullable=False) diff --git a/app/routes.py b/app/routes.py index 9b1e739..8d4ed96 100644 --- a/app/routes.py +++ b/app/routes.py @@ -19,8 +19,19 @@ async def log_callback_request( callback_data: CallbackRequest ): """记录回调请求到数据库""" + # 获取客户端IP地址 + client_ip = request.client.host if request.client else None + + # 获取服务器IP地址 + server_ip = None + if hasattr(request, 'scope') and 'server' in request.scope: + server_host, server_port = request.scope['server'] + server_ip = server_host + callback_log = CallbackLog( site_id=site_id, # 记录siteId + remote_address=client_ip, + server_ip=server_ip, request_url=str(request.url), request_headers=dict(request.headers), request_body=callback_data.model_dump() diff --git a/test_callback.py b/test_callback.py index 9745da1..0443430 100644 --- a/test_callback.py +++ b/test_callback.py @@ -297,6 +297,50 @@ class TestAITalkCallback: assert response.status_code == 500 assert "服务器内部错误" in response.json()["detail"] + @patch('app.routes.get_db') + @patch('app.routes.redis_manager') + def test_callback_ip_fields_saved( + self, + mock_redis_manager, + mock_get_db, + sample_callback_request + ): + """测试remote_address和server_ip字段是否正确保存""" + from app.database import CallbackLog + + # 设置模拟对象 + mock_db = AsyncMock() + mock_get_db.return_value = mock_db + + # 模拟Redis锁 + mock_lock = AsyncMock() + mock_lock.__aenter__ = AsyncMock(return_value=None) + mock_lock.__aexit__ = AsyncMock(return_value=None) + mock_redis_manager.create_lock.return_value = mock_lock + + # 模拟外部API调用成功 + with patch('app.routes.call_external_api_with_retry') as mock_api_call: + mock_api_call.return_value = (True, 0) + + response = client.post( + "/ai-talk/callback/test-site-ip-fields", + json=sample_callback_request.model_dump() + ) + + assert response.status_code == 200 + + # 验证log_callback_request被调用 + assert mock_db.add.called + assert mock_db.commit.called + + # 获取传递给add的CallbackLog对象 + call_args = mock_db.add.call_args[0][0] + assert isinstance(call_args, CallbackLog) + + # 验证新字段存在(可能为None,因为测试环境) + assert hasattr(call_args, 'remote_address') + assert hasattr(call_args, 'server_ip') + if __name__ == "__main__": pytest.main([__file__, "-v"]) \ No newline at end of file