897 lines
34 KiB
Python
897 lines
34 KiB
Python
"""API endpoint tests"""
|
|
|
|
import pytest
|
|
from datetime import datetime
|
|
from unittest.mock import patch, AsyncMock
|
|
from httpx import AsyncClient, ASGITransport
|
|
|
|
from app.main import app
|
|
from app.core.config import settings
|
|
from app.core.security import create_access_token
|
|
from app.db.session import get_db
|
|
from app.models.user import User
|
|
from app.schemas.ai import (
|
|
AIProviderStatusResponse,
|
|
PlaygroundSessionResponse,
|
|
PlaygroundSessionState,
|
|
SituationalAnalysisResponse,
|
|
)
|
|
|
|
|
|
@pytest.fixture
|
|
def auth_headers():
|
|
"""Create authentication headers"""
|
|
token = create_access_token({"sub": "1", "username": "testuser"})
|
|
return {"Authorization": f"Bearer {token}"}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_health_check():
|
|
"""Test health check endpoint"""
|
|
transport = ASGITransport(app=app)
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get("/health")
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["status"] == "healthy"
|
|
assert "version" in data
|
|
assert response.headers["x-request-id"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_root_endpoint():
|
|
"""Test root endpoint"""
|
|
transport = ASGITransport(app=app)
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get("/")
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["name"] == settings.PROJECT_NAME
|
|
assert data["version"] == settings.VERSION
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dashboard_stats_without_auth():
|
|
"""Test dashboard stats requires authentication"""
|
|
transport = ASGITransport(app=app)
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get("/api/v1/dashboard/stats")
|
|
assert response.status_code == 401
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dashboard_stats_with_auth(auth_headers):
|
|
"""Test dashboard stats with authentication"""
|
|
with patch("app.api.v1.dashboard.cache.get", return_value=None):
|
|
with patch("app.api.v1.dashboard.cache.set", return_value=True):
|
|
with patch("app.db.session.get_db") as mock_get_db:
|
|
mock_session = AsyncMock()
|
|
mock_result = AsyncMock()
|
|
mock_result.scalar.return_value = 0
|
|
mock_result.fetchall.return_value = []
|
|
mock_session.execute.return_value = mock_result
|
|
|
|
async def mock_db_context():
|
|
yield mock_session
|
|
|
|
mock_get_db.return_value = mock_db_context()
|
|
|
|
transport = ASGITransport(app=app)
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get(
|
|
"/api/v1/dashboard/stats",
|
|
headers=auth_headers,
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert "total_datasources" in data
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_alerts_without_auth():
|
|
"""Test alerts endpoint requires authentication"""
|
|
transport = ASGITransport(app=app)
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get("/api/v1/alerts")
|
|
assert response.status_code == 401
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_datasource_task_status_without_auth():
|
|
"""Test datasource task-status endpoint requires authentication"""
|
|
transport = ASGITransport(app=app)
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get("/api/v1/datasources/1/task-status")
|
|
assert response.status_code == 401
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_alerts_endpoint_with_auth(auth_headers):
|
|
"""Test alerts endpoint with authentication"""
|
|
class _ScalarResult:
|
|
def __init__(self, rows=None, scalar_value=0):
|
|
self._rows = rows or []
|
|
self._scalar_value = scalar_value
|
|
|
|
def scalars(self):
|
|
class _Scalars:
|
|
def __init__(self, rows):
|
|
self._rows = rows
|
|
|
|
def all(self):
|
|
return self._rows
|
|
|
|
return _Scalars(self._rows)
|
|
|
|
def scalar(self):
|
|
return self._scalar_value
|
|
|
|
class _FakeAlertsSession:
|
|
def __init__(self):
|
|
self.calls = 0
|
|
|
|
async def execute(self, _query):
|
|
self.calls += 1
|
|
if self.calls == 1:
|
|
return _ScalarResult(rows=[])
|
|
return _ScalarResult(rows=[], scalar_value=0)
|
|
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="testuser",
|
|
email="test@example.com",
|
|
password_hash="hashed",
|
|
role="admin",
|
|
is_active=True,
|
|
)
|
|
|
|
async def override_get_db():
|
|
yield _FakeAlertsSession()
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
get_db: override_get_db,
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get("/api/v1/alerts", headers=auth_headers)
|
|
assert response.status_code == 200
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_system_log_sources_requires_super_admin(auth_headers):
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="testuser",
|
|
email="test@example.com",
|
|
password_hash="hashed",
|
|
role="admin",
|
|
is_active=True,
|
|
)
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get("/api/v1/system/logs/sources", headers=auth_headers)
|
|
assert response.status_code == 403
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_system_log_sources_with_super_admin(auth_headers):
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="root",
|
|
email="root@example.com",
|
|
password_hash="hashed",
|
|
role="super_admin",
|
|
is_active=True,
|
|
)
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
with patch(
|
|
"app.api.v1.system_control.list_log_sources",
|
|
return_value=[
|
|
{
|
|
"source_id": "backend",
|
|
"name": "后端服务",
|
|
"kind": "file",
|
|
"location": "/tmp/planet_backend.log",
|
|
"description": "FastAPI 后端、调度器和采集任务共享日志。",
|
|
"category": "service",
|
|
"status": "ok",
|
|
}
|
|
],
|
|
):
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get("/api/v1/system/logs/sources", headers=auth_headers)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["items"][0]["source_id"] == "backend"
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_system_log_snapshot_with_super_admin(auth_headers):
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="root",
|
|
email="root@example.com",
|
|
password_hash="hashed",
|
|
role="super_admin",
|
|
is_active=True,
|
|
)
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
with patch(
|
|
"app.api.v1.system_control.read_log_snapshot",
|
|
return_value={
|
|
"source_id": "backend",
|
|
"name": "后端服务",
|
|
"kind": "file",
|
|
"location": "/tmp/planet_backend.log",
|
|
"description": "FastAPI 后端、调度器和采集任务共享日志。",
|
|
"category": "service",
|
|
"status": "ok",
|
|
"level": "all",
|
|
"selected_levels": [],
|
|
"search_query": "",
|
|
"available_levels": ["all", "error", "warning", "info", "debug"],
|
|
"daily_markers": [],
|
|
"line_limit": 50,
|
|
"line_count": 2,
|
|
"lines": ["line 1", "line 2"],
|
|
},
|
|
):
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get("/api/v1/system/logs/backend?limit=50", headers=auth_headers)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["source_id"] == "backend"
|
|
assert data["line_count"] == 2
|
|
assert data["lines"] == ["line 1", "line 2"]
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_system_log_snapshot_supports_level_filter(auth_headers):
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="root",
|
|
email="root@example.com",
|
|
password_hash="hashed",
|
|
role="super_admin",
|
|
is_active=True,
|
|
)
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
with patch(
|
|
"app.api.v1.system_control.read_log_snapshot",
|
|
return_value={
|
|
"source_id": "backend",
|
|
"name": "后端服务",
|
|
"kind": "file",
|
|
"location": "/tmp/planet_backend.log",
|
|
"description": "FastAPI 后端、调度器和采集任务共享日志。",
|
|
"category": "service",
|
|
"status": "ok",
|
|
"level": "error",
|
|
"selected_levels": ["error"],
|
|
"search_query": "",
|
|
"available_levels": ["all", "error", "warning", "info", "debug"],
|
|
"daily_markers": [],
|
|
"line_limit": 50,
|
|
"line_count": 1,
|
|
"lines": ["ERROR: failed"],
|
|
},
|
|
):
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get("/api/v1/system/logs/backend?limit=50&level=error", headers=auth_headers)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["level"] == "error"
|
|
assert data["lines"] == ["ERROR: failed"]
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_system_log_snapshot_supports_date_range_filter(auth_headers):
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="root",
|
|
email="root@example.com",
|
|
password_hash="hashed",
|
|
role="super_admin",
|
|
is_active=True,
|
|
)
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
with patch(
|
|
"app.api.v1.system_control.read_log_snapshot",
|
|
return_value={
|
|
"source_id": "backend",
|
|
"name": "后端服务",
|
|
"kind": "file",
|
|
"location": "/tmp/planet_backend.log",
|
|
"description": "FastAPI 后端、调度器和采集任务共享日志。",
|
|
"category": "service",
|
|
"status": "ok",
|
|
"level": "all",
|
|
"selected_levels": [],
|
|
"search_query": "",
|
|
"available_levels": ["all", "error", "warning", "info", "debug"],
|
|
"daily_markers": [],
|
|
"line_limit": 50,
|
|
"line_count": 1,
|
|
"lines": ["2026-04-23 INFO: service started"],
|
|
},
|
|
) as mock_read_log_snapshot:
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get(
|
|
"/api/v1/system/logs/backend?limit=50&start_date=2026-04-20&end_date=2026-04-23",
|
|
headers=auth_headers,
|
|
)
|
|
assert response.status_code == 200
|
|
mock_read_log_snapshot.assert_called_once_with(
|
|
"backend",
|
|
50,
|
|
level="all",
|
|
levels=None,
|
|
start_date="2026-04-20",
|
|
end_date="2026-04-23",
|
|
search=None,
|
|
)
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_system_log_snapshot_supports_levels_and_search_filter(auth_headers):
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="root",
|
|
email="root@example.com",
|
|
password_hash="hashed",
|
|
role="super_admin",
|
|
is_active=True,
|
|
)
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
with patch(
|
|
"app.api.v1.system_control.read_log_snapshot",
|
|
return_value={
|
|
"source_id": "backend",
|
|
"name": "后端服务",
|
|
"kind": "file",
|
|
"location": "/tmp/planet_backend.log",
|
|
"description": "FastAPI 后端、调度器和采集任务共享日志。",
|
|
"category": "service",
|
|
"status": "ok",
|
|
"level": "all",
|
|
"selected_levels": ["error", "warning"],
|
|
"search_query": "timeout",
|
|
"available_levels": ["all", "error", "warning", "info", "debug"],
|
|
"daily_markers": [],
|
|
"line_limit": 50,
|
|
"line_count": 1,
|
|
"lines": ["2026-04-23 10:00:00 ERROR timeout"],
|
|
},
|
|
) as mock_read_log_snapshot:
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get(
|
|
"/api/v1/system/logs/backend?limit=50&levels=error,warning&search=timeout",
|
|
headers=auth_headers,
|
|
)
|
|
assert response.status_code == 200
|
|
mock_read_log_snapshot.assert_called_once_with(
|
|
"backend",
|
|
50,
|
|
level="all",
|
|
levels="error,warning",
|
|
start_date=None,
|
|
end_date=None,
|
|
search="timeout",
|
|
)
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_system_log_snapshot_rejects_invalid_date_range(auth_headers):
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="root",
|
|
email="root@example.com",
|
|
password_hash="hashed",
|
|
role="super_admin",
|
|
is_active=True,
|
|
)
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get(
|
|
"/api/v1/system/logs/backend?start_date=2026-04-31",
|
|
headers=auth_headers,
|
|
)
|
|
assert response.status_code == 400
|
|
assert "start_date must be in YYYY-MM-DD format" in response.json()["detail"]
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_ingest_earth_client_log_accepts_public_events():
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
with patch("app.api.v1.system_control.append_buffer_log") as mock_append_buffer_log:
|
|
with patch("app.api.v1.system_control.record_system_log", new_callable=AsyncMock) as mock_record_system_log:
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.post(
|
|
"/api/v1/system/logs/earth-client",
|
|
json={
|
|
"level": "error",
|
|
"message": "登陆点加载失败: 登陆点接口返回 HTTP 500",
|
|
"category": "startup-load",
|
|
"module": "layer-startup",
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["accepted"] is True
|
|
assert data["source_id"] == "earth-client"
|
|
mock_append_buffer_log.assert_called_once()
|
|
mock_record_system_log.assert_awaited_once()
|
|
persisted_kwargs = mock_record_system_log.await_args.kwargs
|
|
assert persisted_kwargs["source"] == "earth-client"
|
|
assert persisted_kwargs["event"] == "earth.client.runtime_log"
|
|
assert persisted_kwargs["category"] == "startup-load"
|
|
assert persisted_kwargs["level"] == "error"
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_request_id_header_is_echoed_when_provided():
|
|
transport = ASGITransport(app=app)
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get("/health", headers={"X-Request-ID": "planet-test-request"})
|
|
assert response.status_code == 200
|
|
assert response.headers["x-request-id"] == "planet-test-request"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_invalid_token():
|
|
"""Test that invalid token is rejected"""
|
|
transport = ASGITransport(app=app)
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get(
|
|
"/api/v1/dashboard/stats",
|
|
headers={"Authorization": "Bearer invalid_token"},
|
|
)
|
|
assert response.status_code == 401
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_ai_provider_status_with_auth(auth_headers):
|
|
"""Test AI provider status endpoint"""
|
|
class _FakeAIProviderClient:
|
|
async def get_status(self, request_id=None):
|
|
return AIProviderStatusResponse(
|
|
provider="minimax",
|
|
api="anthropic-messages",
|
|
enabled=True,
|
|
configured=True,
|
|
model="test-model",
|
|
base_url="http://aiprovider:8010",
|
|
)
|
|
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="testuser",
|
|
email="test@example.com",
|
|
password_hash="hashed",
|
|
role="admin",
|
|
is_active=True,
|
|
)
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
__import__("app.services.ai_client", fromlist=["get_ai_provider_client"]).get_ai_provider_client: lambda: _FakeAIProviderClient(),
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get("/api/v1/ai/provider/status", headers=auth_headers)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert "provider" in data
|
|
assert "api" in data
|
|
assert "configured" in data
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_ai_situational_analysis_returns_503_when_disabled(auth_headers):
|
|
"""Test AI analysis endpoint proxies provider service response"""
|
|
class _FakeAIProviderClient:
|
|
async def analyze(self, _payload, request_id=None):
|
|
return SituationalAnalysisResponse(
|
|
provider="openai_compatible",
|
|
model="test-model",
|
|
content="1) 态势摘要: 测试返回",
|
|
content_blocks=[],
|
|
text_blocks=["1) 态势摘要: 测试返回"],
|
|
thinking_blocks=[],
|
|
raw_response={"id": "mock-response"},
|
|
)
|
|
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="testuser",
|
|
email="test@example.com",
|
|
password_hash="hashed",
|
|
role="admin",
|
|
is_active=True,
|
|
)
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
__import__("app.services.ai_client", fromlist=["get_ai_provider_client"]).get_ai_provider_client: lambda: _FakeAIProviderClient(),
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.post(
|
|
"/api/v1/ai/situational-awareness/analyze",
|
|
headers=auth_headers,
|
|
json={
|
|
"title": "BGP 异常研判",
|
|
"objective": "给出当前异常的风险摘要和建议动作",
|
|
"observations": ["collector A 在 5 分钟内出现多个 origin 变更"],
|
|
"constraints": ["不要假设缺失数据"],
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["provider"] == "openai_compatible"
|
|
assert data["content"]
|
|
assert "content_blocks" in data
|
|
assert "text_blocks" in data
|
|
assert "thinking_blocks" in data
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_playground_session_with_auth(auth_headers):
|
|
"""Test playground session restore endpoint."""
|
|
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="testuser",
|
|
email="test@example.com",
|
|
password_hash="hashed",
|
|
role="admin",
|
|
is_active=True,
|
|
)
|
|
|
|
async def override_get_db():
|
|
yield AsyncMock()
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
get_db: override_get_db,
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
with patch(
|
|
"app.api.v1.ai.get_playground_session",
|
|
new=AsyncMock(
|
|
return_value=PlaygroundSessionResponse(
|
|
id="1",
|
|
session_key="default",
|
|
title="Playground 会话",
|
|
state=PlaygroundSessionState(
|
|
messages=[{"id": "msg-1", "role": "user", "content": "hello"}],
|
|
title="测试标题",
|
|
objective="测试目标",
|
|
),
|
|
created_at="2026-04-10T00:00:00+00:00",
|
|
updated_at="2026-04-10T00:00:00+00:00",
|
|
)
|
|
),
|
|
):
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.get("/api/v1/ai/playground/session", headers=auth_headers)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["session_key"] == "default"
|
|
assert data["state"]["messages"][0]["content"] == "hello"
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_save_playground_session_with_auth(auth_headers):
|
|
"""Test playground session save endpoint."""
|
|
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="testuser",
|
|
email="test@example.com",
|
|
password_hash="hashed",
|
|
role="admin",
|
|
is_active=True,
|
|
)
|
|
|
|
async def override_get_db():
|
|
yield AsyncMock()
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
get_db: override_get_db,
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
with patch(
|
|
"app.api.v1.ai.upsert_playground_session",
|
|
new=AsyncMock(
|
|
return_value=PlaygroundSessionResponse(
|
|
id="1",
|
|
session_key="default",
|
|
title="测试标题",
|
|
state=PlaygroundSessionState(
|
|
messages=[{"id": "msg-1", "role": "user", "content": "hello"}],
|
|
title="测试标题",
|
|
objective="测试目标",
|
|
),
|
|
created_at="2026-04-10T00:00:00+00:00",
|
|
updated_at="2026-04-10T00:00:00+00:00",
|
|
)
|
|
),
|
|
):
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.put(
|
|
"/api/v1/ai/playground/session",
|
|
headers=auth_headers,
|
|
json={
|
|
"session_key": "default",
|
|
"title": "测试标题",
|
|
"state": {
|
|
"messages": [{"id": "msg-1", "role": "user", "content": "hello"}],
|
|
"selectedPresetKey": "bgp-brief",
|
|
"title": "测试标题",
|
|
"objective": "测试目标",
|
|
"constraints": "",
|
|
"inputValue": "",
|
|
"analysis": None,
|
|
"latestAnalysisMessageId": None,
|
|
"analysisMeta": {},
|
|
"helpExpanded": True,
|
|
},
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["title"] == "测试标题"
|
|
assert data["state"]["objective"] == "测试目标"
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_ai_bgp_brief_endpoint_persists_fact_snapshot(auth_headers):
|
|
class _FakeAIProviderClient:
|
|
async def analyze(self, _payload, request_id=None):
|
|
return SituationalAnalysisResponse(
|
|
provider="minimax",
|
|
model="MiniMax-M2.5",
|
|
content="# BGP AI 简报\n\n事实摘要:测试",
|
|
content_blocks=[],
|
|
text_blocks=["# BGP AI 简报\n\n事实摘要:测试"],
|
|
thinking_blocks=[],
|
|
raw_response={"id": "mock-bgp-brief"},
|
|
)
|
|
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="testuser",
|
|
email="test@example.com",
|
|
password_hash="hashed",
|
|
role="admin",
|
|
is_active=True,
|
|
)
|
|
|
|
async def override_get_db():
|
|
yield AsyncMock()
|
|
|
|
async def _fake_build_bgp_brief_request(_db, **_kwargs):
|
|
request_payload = __import__("app.schemas.ai", fromlist=["SituationalAnalysisRequest"]).SituationalAnalysisRequest(
|
|
title="BGP 态势 AI 简报",
|
|
objective="生成值班简报",
|
|
observations=["事实A", "事实B"],
|
|
constraints=["不要编造"],
|
|
context={"incident_total": 2, "active_collectors": 3},
|
|
)
|
|
return request_payload, ["事实A", "事实B"], {"incident_total": 2, "active_collectors": 3}
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
__import__("app.services.ai_client", fromlist=["get_ai_provider_client"]).get_ai_provider_client: lambda: _FakeAIProviderClient(),
|
|
get_db: override_get_db,
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
with patch("app.api.v1.ai.build_bgp_brief_request", side_effect=_fake_build_bgp_brief_request):
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.post("/api/v1/ai/bgp/brief", headers=auth_headers, json={})
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["facts"] == ["事实A", "事实B"]
|
|
assert data["context"]["incident_total"] == 2
|
|
assert data["content_markdown"]
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_ai_alert_brief_endpoint_with_auth(auth_headers):
|
|
class _FakeAIProviderClient:
|
|
async def analyze(self, _payload, request_id=None):
|
|
return SituationalAnalysisResponse(
|
|
provider="minimax",
|
|
model="MiniMax-M2.7",
|
|
content="事实摘要:告警测试。风险研判:告警测试。建议动作:告警测试。",
|
|
content_blocks=[],
|
|
text_blocks=["事实摘要:告警测试。风险研判:告警测试。建议动作:告警测试。"],
|
|
thinking_blocks=[],
|
|
raw_response={"id": "mock-alert-brief"},
|
|
)
|
|
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="testuser",
|
|
email="test@example.com",
|
|
password_hash="hashed",
|
|
role="admin",
|
|
is_active=True,
|
|
)
|
|
|
|
async def override_get_db():
|
|
yield AsyncMock()
|
|
|
|
async def _fake_build_alert_brief_request(_db, **_kwargs):
|
|
request_payload = __import__("app.schemas.ai", fromlist=["SituationalAnalysisRequest"]).SituationalAnalysisRequest(
|
|
title="告警态势 AI 简报",
|
|
objective="输出告警简报",
|
|
observations=["告警事实A", "告警事实B"],
|
|
constraints=["不要编造"],
|
|
context={"active_alerts": 3, "top_datasources": {"bgp": 2}},
|
|
)
|
|
return request_payload, ["告警事实A", "告警事实B"], {"active_alerts": 3, "top_datasources": {"bgp": 2}}
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
__import__("app.services.ai_client", fromlist=["get_ai_provider_client"]).get_ai_provider_client: lambda: _FakeAIProviderClient(),
|
|
get_db: override_get_db,
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
with patch("app.api.v1.ai.build_alert_brief_request", side_effect=_fake_build_alert_brief_request):
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.post("/api/v1/ai/alerts/brief", headers=auth_headers, json={})
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["title"] == "告警态势 AI 简报"
|
|
assert data["facts"] == ["告警事实A", "告警事实B"]
|
|
assert data["context"]["active_alerts"] == 3
|
|
assert data["content"]
|
|
finally:
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_ai_situational_alert_brief_endpoint_with_auth(auth_headers):
|
|
class _FakeAIProviderClient:
|
|
async def analyze(self, _payload, request_id=None):
|
|
return SituationalAnalysisResponse(
|
|
provider="minimax",
|
|
model="MiniMax-M2.7",
|
|
content="事实摘要:态势测试。风险研判:态势测试。建议动作:态势测试。",
|
|
content_blocks=[],
|
|
text_blocks=["事实摘要:态势测试。风险研判:态势测试。建议动作:态势测试。"],
|
|
thinking_blocks=[],
|
|
raw_response={"id": "mock-situational-brief"},
|
|
)
|
|
|
|
def override_get_current_user():
|
|
return User(
|
|
id=1,
|
|
username="testuser",
|
|
email="test@example.com",
|
|
password_hash="hashed",
|
|
role="admin",
|
|
is_active=True,
|
|
)
|
|
|
|
async def override_get_db():
|
|
yield AsyncMock()
|
|
|
|
async def _fake_build_situational_alert_brief_request(_db):
|
|
request_payload = __import__("app.schemas.ai", fromlist=["SituationalAnalysisRequest"]).SituationalAnalysisRequest(
|
|
title="态势告警 AI 简报",
|
|
objective="输出态势告警简报",
|
|
observations=["态势事实A", "态势事实B"],
|
|
constraints=["不要编造"],
|
|
context={"active_system_alerts": 2, "active_bgp_incidents": 1},
|
|
)
|
|
return request_payload, ["态势事实A", "态势事实B"], {"active_system_alerts": 2, "active_bgp_incidents": 1}
|
|
|
|
app.dependency_overrides = {
|
|
__import__("app.core.security", fromlist=["get_current_user"]).get_current_user: override_get_current_user,
|
|
__import__("app.services.ai_client", fromlist=["get_ai_provider_client"]).get_ai_provider_client: lambda: _FakeAIProviderClient(),
|
|
get_db: override_get_db,
|
|
}
|
|
transport = ASGITransport(app=app)
|
|
try:
|
|
with patch("app.api.v1.ai.build_situational_alert_brief_request", side_effect=_fake_build_situational_alert_brief_request):
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
|
response = await client.post("/api/v1/ai/situational-alerts/brief", headers=auth_headers, json={})
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["title"] == "态势告警 AI 简报"
|
|
assert data["facts"] == ["态势事实A", "态势事实B"]
|
|
assert data["context"]["active_system_alerts"] == 2
|
|
assert data["content"]
|
|
finally:
|
|
app.dependency_overrides.clear()
|