增加remote_address和server_ip

This commit is contained in:
mark.tian
2025-12-02 18:31:51 +08:00
parent 9cb45204c9
commit 49afdc9692
3 changed files with 57 additions and 0 deletions

View File

@@ -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)

View File

@@ -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()

View File

@@ -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"])