增加remote_address和server_ip
This commit is contained in:
@@ -14,6 +14,8 @@ class CallbackLog(Base):
|
|||||||
|
|
||||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||||
site_id = Column(String(100), nullable=False, comment="站点ID") # 新增siteId字段
|
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_url = Column(String(500), nullable=False)
|
||||||
request_headers = Column(JSON, nullable=False)
|
request_headers = Column(JSON, nullable=False)
|
||||||
request_body = Column(JSON, nullable=False)
|
request_body = Column(JSON, nullable=False)
|
||||||
|
|||||||
@@ -19,8 +19,19 @@ async def log_callback_request(
|
|||||||
callback_data: CallbackRequest
|
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(
|
callback_log = CallbackLog(
|
||||||
site_id=site_id, # 记录siteId
|
site_id=site_id, # 记录siteId
|
||||||
|
remote_address=client_ip,
|
||||||
|
server_ip=server_ip,
|
||||||
request_url=str(request.url),
|
request_url=str(request.url),
|
||||||
request_headers=dict(request.headers),
|
request_headers=dict(request.headers),
|
||||||
request_body=callback_data.model_dump()
|
request_body=callback_data.model_dump()
|
||||||
|
|||||||
@@ -297,6 +297,50 @@ class TestAITalkCallback:
|
|||||||
assert response.status_code == 500
|
assert response.status_code == 500
|
||||||
assert "服务器内部错误" in response.json()["detail"]
|
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__":
|
if __name__ == "__main__":
|
||||||
pytest.main([__file__, "-v"])
|
pytest.main([__file__, "-v"])
|
||||||
Reference in New Issue
Block a user