125 lines
4.7 KiB
Python
125 lines
4.7 KiB
Python
"""REST API 测试(标准库 http.server)。"""
|
||
import _bootstrap # noqa: F401
|
||
import json
|
||
import threading
|
||
import unittest
|
||
from http.server import ThreadingHTTPServer
|
||
from urllib.error import HTTPError
|
||
from urllib.request import Request, urlopen
|
||
|
||
from fault_log_analyzer.api import create_handler
|
||
from fault_log_analyzer.models import FaultLog, RootCause
|
||
from fault_log_analyzer.storage import MemoryStorage
|
||
|
||
|
||
class ApiTestBase(unittest.TestCase):
|
||
def setUp(self):
|
||
self.storage = MemoryStorage()
|
||
handler = create_handler(self.storage)
|
||
self.server = ThreadingHTTPServer(("127.0.0.1", 0), handler)
|
||
self.base = f"http://127.0.0.1:{self.server.server_address[1]}"
|
||
self.thread = threading.Thread(target=self.server.serve_forever, daemon=True)
|
||
self.thread.start()
|
||
|
||
def tearDown(self):
|
||
self.server.shutdown()
|
||
self.server.server_close()
|
||
self.thread.join(timeout=2)
|
||
|
||
def request(self, method, path, body=None):
|
||
url = self.base + path
|
||
data = None
|
||
headers = {}
|
||
if body is not None:
|
||
data = json.dumps(body).encode("utf-8")
|
||
headers["Content-Type"] = "application/json"
|
||
req = Request(url, data=data, method=method, headers=headers)
|
||
try:
|
||
resp = urlopen(req, timeout=5)
|
||
return resp.status, json.loads(resp.read().decode("utf-8"))
|
||
except HTTPError as e:
|
||
return e.code, json.loads(e.read().decode("utf-8"))
|
||
|
||
|
||
class TestHealth(ApiTestBase):
|
||
def test_healthz(self):
|
||
status, body = self.request("GET", "/healthz")
|
||
self.assertEqual(status, 200)
|
||
self.assertEqual(body["status"], "ok")
|
||
|
||
|
||
class TestFaultLogs(ApiTestBase):
|
||
def test_list_empty(self):
|
||
status, body = self.request("GET", "/api/v1/fault-logs")
|
||
self.assertEqual(status, 200)
|
||
self.assertEqual(body["data"]["total"], 0)
|
||
|
||
def test_list_and_detail(self):
|
||
self.storage.save_fault_log(
|
||
FaultLog(fault_log_id="fl-1", host_id="h-1", message="no space", level="ERROR", fault_type="disk_full")
|
||
)
|
||
status, body = self.request("GET", "/api/v1/fault-logs")
|
||
self.assertEqual(body["data"]["total"], 1)
|
||
status, body = self.request("GET", "/api/v1/fault-logs/fl-1")
|
||
self.assertEqual(body["data"]["fault_log_id"], "fl-1")
|
||
|
||
def test_detail_not_found(self):
|
||
status, body = self.request("GET", "/api/v1/fault-logs/nope")
|
||
self.assertEqual(status, 404)
|
||
self.assertEqual(body["code"], 40401)
|
||
|
||
def test_invalid_page_param(self):
|
||
status, body = self.request("GET", "/api/v1/fault-logs?page=abc")
|
||
self.assertEqual(status, 400)
|
||
self.assertEqual(body["code"], 40001)
|
||
|
||
def test_root_cause(self):
|
||
self.storage.save_fault_log(FaultLog(fault_log_id="fl-1", host_id="h-1", message="x"))
|
||
self.storage.save_root_cause(RootCause(fault_log_id="fl-1", cause_type="disk_full"))
|
||
status, body = self.request("GET", "/api/v1/fault-logs/fl-1/root-cause")
|
||
self.assertEqual(body["data"]["cause_type"], "disk_full")
|
||
|
||
def test_root_cause_not_found(self):
|
||
status, body = self.request("GET", "/api/v1/fault-logs/nope/root-cause")
|
||
self.assertEqual(status, 404)
|
||
|
||
|
||
class TestFaultTypes(ApiTestBase):
|
||
def test_list_and_create(self):
|
||
status, body = self.request("GET", "/api/v1/fault-types")
|
||
self.assertEqual(body["data"]["items"], [])
|
||
status, body = self.request("POST", "/api/v1/fault-types", {"fault_type": "x", "name": "X"})
|
||
self.assertEqual(status, 201)
|
||
self.assertEqual(body["data"]["fault_type"], "x")
|
||
status, body = self.request("GET", "/api/v1/fault-types")
|
||
self.assertEqual(len(body["data"]["items"]), 1)
|
||
|
||
def test_create_missing_fields(self):
|
||
status, body = self.request("POST", "/api/v1/fault-types", {"name": "X"})
|
||
self.assertEqual(status, 400)
|
||
self.assertEqual(body["code"], 40002)
|
||
|
||
|
||
class TestFaultFilters(ApiTestBase):
|
||
def test_list_and_create(self):
|
||
status, body = self.request("GET", "/api/v1/fault-filters")
|
||
self.assertEqual(body["data"]["items"], [])
|
||
status, body = self.request("POST", "/api/v1/fault-filters", {"name": "r1", "level": "ERROR"})
|
||
self.assertEqual(status, 201)
|
||
self.assertEqual(body["data"]["name"], "r1")
|
||
|
||
def test_create_missing_name(self):
|
||
status, body = self.request("POST", "/api/v1/fault-filters", {"level": "ERROR"})
|
||
self.assertEqual(status, 400)
|
||
self.assertEqual(body["code"], 40002)
|
||
|
||
|
||
class TestNotFound(ApiTestBase):
|
||
def test_unknown_route(self):
|
||
status, body = self.request("GET", "/api/v1/unknown")
|
||
self.assertEqual(status, 404)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
unittest.main()
|