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