Files
truf-server/tests/test_zai_provider.py
T
2026-09-30 20:30:56 +03:00

181 lines
7.7 KiB
Python

from pathlib import Path
import re
import sys
import unittest
from unittest import mock
import yaml
ROOT = Path(__file__).resolve().parents[1]
APP_DIR = ROOT / 'app'
sys.path.insert(0, str(APP_DIR))
import dashboard
import keycheck_runner
from keycheck_candidates import extract_candidates, extract_structured_candidates
from keycheckers.zai import zaiKeycheck as zai
def response(status, payload):
item = mock.Mock(status_code=status, text='fixture response')
item.json.return_value = payload
return item
class ZaiCandidateTests(unittest.TestCase):
def test_custom_detector_accepts_the_observed_dotted_shape(self):
policy = yaml.safe_load(
(APP_DIR / 'trufflehog-custom-detectors.yaml').read_text(encoding='utf-8')
)
detector = next(item for item in policy['detectors'] if item['name'] == 'ZaiGLM')
patterns = [re.compile(value) for value in detector['regex'].values()]
key = ('f' * 32) + '.' + ('G' * 16)
self.assertTrue(any(pattern.search(f'ZAI_API_KEY={key}') for pattern in patterns))
self.assertFalse(any(
pattern.search(f'ZAI_API_KEY={("f" * 32)}.{("G" * 15)}')
for pattern in patterns
))
def test_real_world_dotted_shape_is_bounded_and_extracted(self):
key = ('a' * 32) + '.' + ('B' * 16)
finding = {
'DetectorName': 'ZaiGLM',
'Raw': key,
'ScannerContext': {'provider_hint': 'zai'},
}
candidates = list(extract_candidates(finding))
self.assertEqual(zai.key_rejection_reason(key), '')
self.assertEqual(len(candidates), 1)
self.assertEqual(candidates[0].service, 'zai')
self.assertEqual(candidates[0].secret_text, key)
self.assertTrue(zai.key_rejection_reason(('a' * 32) + '.' + ('B' * 15)))
def test_structured_zai_environment_key_is_extracted(self):
key = ('b' * 32) + '.' + ('C' * 16)
candidates = list(extract_structured_candidates({
'origin': 'fixture',
'contexts': [{'key': 'ZAI_API_KEY', 'value': key, 'path': 'env.0'}],
}))
self.assertEqual([(candidate.service, candidate.secret_text) for candidate in candidates], [('zai', key)])
class ZaiKeycheckTests(unittest.TestCase):
def setUp(self):
self.key = ('c' * 32) + '.' + ('D' * 16)
def test_generation_probe_is_required_for_valid_status(self):
payload = {'data': [{'id': 'glm-5'}, {'id': 'glm-4.7'}]}
completion = response(200, {'choices': [{'message': {'content': 'p'}}]})
with mock.patch.object(zai.requests, 'get', return_value=response(200, payload)) as request, \
mock.patch.object(zai.requests, 'post', return_value=completion) as probe:
result = zai.check_base_url(self.key, zai.DEFAULT_BASE_URLS[0], None, 5)
self.assertEqual(result['status'], 'VALID')
self.assertTrue(result['authenticated'])
self.assertEqual(result['model_count'], 2)
self.assertEqual(result['llm_probe_status'], 'GENERATION_OK')
self.assertEqual(result['llm_probe_model'], 'glm-5.2')
self.assertTrue(request.call_args.args[0].endswith('/models'))
self.assertTrue(probe.call_args.args[0].endswith('/chat/completions'))
self.assertEqual(probe.call_args.kwargs['json']['model'], 'glm-5.2')
self.assertEqual(probe.call_args.kwargs['json']['max_tokens'], 1)
def test_authenticated_probe_failures_never_remain_valid(self):
models = response(200, {'data': [{'id': 'glm-4.7'}]})
cases = (
('NO_BALANCE', response(429, {
'error': {'code': '1113', 'message': 'Insufficient balance'},
})),
('LIMITED', response(429, {
'error': {'code': '1302', 'message': 'Usage limit reached'},
})),
('RESTRICTED', response(403, {
'error': {'code': '1220', 'message': 'Model access denied'},
})),
('UNKNOWN', response(200, {'choices': []})),
)
for expected, probe_response in cases:
with self.subTest(expected=expected), \
mock.patch.object(zai.requests, 'get', return_value=models), \
mock.patch.object(zai.requests, 'post', return_value=probe_response):
result = zai.check_base_url(self.key, zai.DEFAULT_BASE_URLS[0], None, 5)
self.assertEqual(result['status'], expected)
self.assertNotEqual(result['status'], 'VALID')
self.assertTrue(result['authenticated'])
self.assertEqual(result['llm_probe_status'], expected)
def test_probe_network_failure_is_authenticated_but_not_alive(self):
models = response(200, {'data': [{'id': 'glm-4.7'}]})
with mock.patch.object(zai.requests, 'get', return_value=models), \
mock.patch.object(zai.requests, 'post', side_effect=zai.requests.Timeout('fixture')):
result = zai.check_base_url(self.key, zai.DEFAULT_BASE_URLS[0], None, 1)
self.assertEqual(result['status'], 'NETWORK')
self.assertTrue(result['authenticated'])
self.assertEqual(result['llm_probe_status'], 'NETWORK')
def test_authentication_business_codes_are_definitive(self):
error = zai.parse_error(response(401, {
'error': {'code': '1000', 'message': 'Authentication Failed'},
}), self.key)
self.assertEqual(zai.classify_error(error), ('DEAD', False))
def test_balance_and_plan_codes_prove_provider_match(self):
self.assertEqual(
zai.classify_error({'http_status': 429, 'code': '1113'}),
('NO_BALANCE', True),
)
self.assertEqual(
zai.classify_error({'http_status': 429, 'code': '1310'}),
('LIMITED', True),
)
self.assertEqual(
zai.classify_error({'http_status': 403, 'code': '1220'}),
('RESTRICTED', True),
)
self.assertEqual(
zai.classify_error({'http_status': 429, 'message': 'Insufficient balance'}),
('NO_BALANCE', True),
)
def test_global_rejection_falls_back_to_china_endpoint(self):
invalid = response(401, {'error': {'code': '1000', 'message': 'Authentication Failed'}})
valid = response(200, {'data': [{'id': 'glm-5'}]})
completion = response(200, {'choices': [{'message': {'content': 'p'}}]})
with mock.patch.object(zai.requests, 'get', side_effect=[invalid, valid]) as request, \
mock.patch.object(zai.requests, 'post', return_value=completion):
result = zai.check_key(self.key, zai.DEFAULT_BASE_URLS, None, 5)
self.assertEqual(request.call_count, 2)
self.assertEqual(result['status'], 'VALID')
self.assertEqual(result['region'], 'open.bigmodel.cn')
def test_network_and_server_failures_remain_retryable(self):
self.assertEqual(
zai.classify_error({'http_status': 500, 'code': '1200'}),
('NETWORK', False),
)
with mock.patch.object(zai.requests, 'get', side_effect=zai.requests.Timeout('fixture')):
result = zai.check_base_url(self.key, zai.DEFAULT_BASE_URLS[0], None, 1)
self.assertEqual(result['status'], 'NETWORK')
self.assertNotIn(self.key, result['message'])
class ZaiRegistryTests(unittest.TestCase):
def test_zai_and_resolver_are_registered(self):
for service in ('zai', 'provider_resolver'):
self.assertIn(service, keycheck_runner.SERVICES)
self.assertIn(service, keycheck_runner.SERVICE_CAPABILITIES)
self.assertIn(service, keycheck_runner.POSTGRES_STATUS_FILE_NAMES)
self.assertIn(service, dashboard.VALIDATION_SERVICES)
if __name__ == '__main__':
unittest.main()