Files
planet/backend/tests/test_api.py

557 lines
21 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
@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_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
@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()
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()