Files
xc-metax-c500-advisor-agent/test_main.py

51 lines
1.8 KiB
Python

import unittest
from main import CARD_MODEL, CARD_VENDOR, ValidationError, analyze
class AnalyzeTests(unittest.TestCase):
def complete_payload(self):
return {
"model_name": "example/model",
"hardware": f"{CARD_VENDOR} {CARD_MODEL}",
"framework": "PyTorch",
"backend": "vendor-backend",
"sdk_version": "provided-by-user",
"driver_version": "provided-by-user",
}
def test_complete_preflight_targets_this_card(self):
result = analyze(self.complete_payload())
self.assertEqual(result["verdict"], "preflight_ready")
self.assertTrue(result["target_matches"])
self.assertEqual(result["official_target"]["model"], CARD_MODEL)
def test_mismatched_target_is_rejected(self):
payload = self.complete_payload()
payload["hardware"] = "其他厂商 其他卡型"
result = analyze(payload)
self.assertEqual(result["verdict"], "target_mismatch")
self.assertFalse(result["target_matches"])
def test_missing_information_is_explicit(self):
result = analyze({})
self.assertEqual(result["verdict"], "information_required")
self.assertIn("sdk_version", result["missing_fields"])
self.assertIn("driver_version", result["missing_fields"])
def test_multicard_and_quantization_risks_are_reported(self):
payload = self.complete_payload()
payload.update({"cards": 2, "precision": "int8"})
result = analyze(payload)
joined = " ".join(result["risks"])
self.assertIn("多卡", joined)
self.assertIn("量化", joined)
def test_invalid_card_count_is_rejected(self):
with self.assertRaises(ValidationError):
analyze({"cards": 0})
if __name__ == "__main__":
unittest.main()